Fast as CHITA: Neural Network Pruning with Combinatorial Optimization

Riade Benbaki, Wenyu Chen, Xiang Meng, Hussein Hazimeh, Natalia Ponomareva, Zhe Zhao, Rahul Mazumder

Introduction

Modern neural networks tend to use a large number of parameters (Devlin et al., 2018; He et al., 2016), which leads to high computational costs during inference. A widely used approach to mitigate inference costs is to prune or sparsify pre-trained networks by removing parameters (Blalock et al., 2020). The goal is to obtain a network with significantly fewer parameters and minimal loss in performance. This makes model storage and deployment cheaper and easier, especially in resource-constrained environments.

Generally speaking, there are two main approaches for neural net pruning: (i) magnitude-based and (ii) impact-based. Magnitude-based heuristic methods (e.g., Hanson & Pratt, 1988; Mozer & Smolensky, 1989; Gordon et al., 2020) use the absolute value of weight to determine its importance and whether or not it should be pruned. Since magnitude alone may not be a perfect proxy for weight relevance, alternatives have been proposed. To this end, impact-based pruning methods (e.g. LeCun et al., 1989; Hassibi & Stork, 1992; Singh & Alistarh, 2020) remove weights based on how much their removal would impact the loss function, often using second-order information on the loss function. Both of these approaches, however, may fall short of considering the joint effect of removing (and updating) multiple weights simultaneously. The recent method CBS (Combinatorial Brain Surgeon) (Yu et al., 2022) is an optimization-based approach that considers the joint effect of multiple weights. The authors show that CBS leads to a boost in the performance of the pruned models. However, CBS can be computationally expensive: it makes use of a local model based on the second-order (Hessian) information of the loss function, which can be prohibitively expensive in terms of runtime and/or memory (e.g., CBS takes hours to prune a network with 4.2 million parameters).

Since the local quadratic model approximates the loss function only in a small neighborhood of the current solution (Singh & Alistarh, 2020), we also propose a multi-stage algorithm that updates the local quadratic model during pruning (but without retraining) and solves a more constrained problem in each stage, going from dense weights to sparse ones. Our experiments show that the resulting pruned models have a notably better accuracy compared to that of our single-stage algorithm and other pruning approaches. Furthermore, when used in the gradual pruning setting (Gale et al., 2019; Singh & Alistarh, 2020; Blalock et al., 2020) where re-training between pruning steps is performed, our pruning framework results in significant performance gains compared to state-of-the-art unstructured pruning methods.

Our contributions can be summarized as follows:

A key workhorse of CHITA is a novel IHT-based algorithm to obtain good solutions to the sparse regression formulation. Exploiting problem structure, we propose methods to accelerate convergence and improve pruning performance, such as a new and efficient stepsize selection scheme and rapidly updating weights on the support. This leads to up to 1000x runtime improvement compared to existing network pruning algorithms.

We show performance improvements across various models and datasets. In particular, CHITA results in a 98% sparse (i.e., 98% of weights in dense model are set to zero) MLPNet with 90% test accuracy (3% reduction in test accuracy compared to the dense model), which is a significant improvement over the previously reported best accuracy (55%) by CBS. As an application of our framework, we use it for gradual pruning and observe notable performance gains against state-of-the-art gradual pruning approaches.

Problem Setup and Related Work

In this section we present the general setup with connections to related work—this lays the foundation for our proposed methods discussed in Section 3.

The loss function at ww is as close as possible to the loss before pruning: L(w)≈L(wˉ)\mathcal{L}(w)\approx\mathcal{L}(\bar{w}).

Similar to LeCun et al. (1989); Hassibi & Stork (1992); Singh & Alistarh (2020), we use a local quadratic approximation of L\mathcal{L} around the pre-trained weight wˉ\bar{w}:

With certain choices of gradient and Hessian approximations g≈∇L(wˉ),H≈∇2L(wˉ)g\approx\nabla\mathcal{L}(\bar{w}),H\approx\nabla^{2}\mathcal{L}(\bar{w}), and ignoring higher-order terms, the loss L\mathcal{L} can be locally approximated by:

Pruning the local approximation Q0(w)Q_{0}(w) of the network can be naturally formulated as an optimization problem to minimize Q0(w)Q_{0}(w) subject to a cardinality constraint, i.e.,

For large networks, solving Problem (3) directly (e.g., using iterative optimization methods) is computationally challenging due to the sheer size of the p×pp\times p matrix HH. In Section 3.1, we present an exact, hessian-free reformulation of Problem (3), which is key to our scalable approach.

2 Related Work

Impact-based pruning dates back to the work of LeCun et al. (1989) where the OBD (Optimal Brain Damage) framework is proposed. This approach, along with subsequent ones (Hassibi & Stork, 1992; Singh & Alistarh, 2020; Yu et al., 2022) make use of local approximation (2). It is usually assumed (but not in our work) that wˉ\bar{w} is a local optimum of the loss function, and therefore g=0g=0 and L(w)≈L(wˉ)+12(w−wˉ)⊤H(w−wˉ)\mathcal{L}(w)\approx\mathcal{L}(\bar{w})+\frac{1}{2}(w-\bar{w})^{\top}H(w-\bar{w}). Using this approximation, OBD (Optimal Brain Damage, LeCun et al. (1989)) searches for a single weight ii to prune with minimal increase of the loss function, while also assuming a diagonal Hessian HH. If the ii-th weigth is pruned (wi=0,wj=wˉj∀j≠iw_{i}=0,w_{j}=\bar{w}_{j}\forall j\neq i), then the loss function increases by δLi=wˉi22∇2L(wˉ)ii\delta\mathcal{L}_{i}=\frac{\bar{w}^{2}_{i}}{2\nabla^{2}\mathcal{L}(\bar{w})}_{ii}. This represents a score for each weight, and is used to prune weights in decreasing order of their score. OBS (Optimal Brain Surgeon, Hassibi & Stork (1992)) extends this by no longer assuming a diagonal Hessian, and using the optimality conditions to update the un-pruned weights. The authors also propose using the empirical Fisher information matrix, as an efficient approximation to the Hessian matrix. Layerwise OBS (Dong et al., 2017) proposes to overcome the computational challenge of computing the full (inverse) Hessian needed in OBS by pruning each layer independently, while Singh & Alistarh (2020) use block-diagonal approximations on the Hessian matrix, which they approximate by the empirical Fisher information matrix on a small subset of the training data (n≪Nn\ll N):

While these approaches explore different ways to make the Hessian computationally tractable, they all rely on the OBD/OBS framework of pruning a single weight, and do not to consider the possible interactions that can arise when pruning multiple weights. To this end, Yu et al. (2022) propose CBS (Combinatorial Brain Surgeon) an algorithm to approximately solve (3). While CBS shows impressive improvements in the accuracy of the pruned model over prior work, it operates with the full dense p×pp\times p Hessian HH. This limits scalability both in compute time and memory utilization, as pp is often in the order of millions and more.

As mentioned earlier, most previous work assumes that the pre-trained weights wˉ\bar{w} is a local optimum of the loss function L\mathcal{L}, and thus take the gradient g=0g=0. However, the gradient of the loss function of a pre-trained neural network may not be zero in practice due to early stopping (or approximate optimization) (Yao et al., 2007). Thus, the WoodTaylor approach (Singh & Alistarh, 2020) proposes to approximate the gradient by the stochastic gradient, using the same samples for estimating the Hessian. Namely,

Generally speaking, one-shot pruning methods (LeCun et al., 1989; Singh & Alistarh, 2020; Yu et al., 2022) can be followed by a few fine-tuning and re-training steps to recover some of the accuracy lost when pruning. Furthermore, recent work has shown that pruning and re-training in a gradual fashion (hence the name, gradual pruning) can lead to big accuracy gains (Han et al., 2015; Gale et al., 2019; Zhu & Gupta, 2018). The work of Singh & Alistarh (2020) further shows that gradual pruning, when used with well-performing one-shot pruning algorithms, can outperform state-of-the-art unstructured pruning methods. In this paper, we focus on the one-shot pruning problem and show how our pruning framework outperforms other one-shot pruning methods (see Section 4.1). We then show that, if applied in the gradual pruning setting, our pruning algorithm outperforms existing approaches (see Section 4.2), establishing new state-of-the-art on MobileNetV1 and ResNet50.

Our Proposed Framework: CHITA

Our formulation is based on a critical observation that the Hessian approximation in (4) has a low-rank structure:

Using observation (6) and the gradient expression (5), we note that problem (3) can be equivalently written in the following Hessian-free form:

2 Our proposed algorithm for problem (8)

We present the core ideas of our proposed algorithm for Problem (8), and discuss additional details in Appendix A.

Our optimization framework relies on the IHT algorithm (Blumensath & Davies, 2009; Bertsimas et al., 2016) that optimizes (8) by simultaneously updating the support and the weights. By leveraging the low-rank structure, we can avoid the computational burden of computing the full Hessian matrix, thus reducing complexity.

The basic version of the IHT algorithm can be slow for problems with a million parameters. To improve the computational performance of our algorithm we propose a new line search scheme. Additionally, we use use an active set strategy and schemes to update the weights on the nonzero weights upon support stabilization. Taken together, we obtain notable improvements in computational efficiency and solution quality over traditional IHT, making it a viable option for network pruning problems at scale.

The IHT algorithm operates by taking a gradient step of size τ\tau from the current iteration, then projects it onto the set of points with a fixed number of non-zero coordinates through hard thresholding. Specifically, for any vector xx, let Ik(x)\mathcal{I}_{k}(x) denote the indices of kk components of xx that have the largest absolute value. The hard thresholding operator Pk(x)P_{k}(x) is defined as yi=xiy_{i}=x_{i} if i∈Ik(x)i\in\mathcal{I}_{k}(x), and yi=0y_{i}=0 if i∉Ik(x)i\notin\mathcal{I}_{k}(x); where yiy_{i} is the ii-th coordinate of Pk(x)P_{k}(x). IHT applied to problem (8) leads to the following update:

where τs>0\tau^{s}>0 is a suitable stepsize. The computation of HT(wt,k,τs)\textsf{HT}(w^{t},k,\tau^{s}) involves only matrix-vector multiplications with AA (or A⊤A^{\top}) and a vector, which has a total computation cost of O(np)O(np). This is a significant reduction compared to the O(p2)O(p^{2}) cost while using the full Hessian matrix as Singh & Alistarh (2020); Yu et al. (2022) do.

Active set strategy. In an effort to further facilitate the efficiency of the IHT method, we propose using an active set strategy, which has been proven successful in various contexts such as (Nocedal & Wright, 1999; Friedman et al., 2010; Hazimeh & Mazumder, 2020). This strategy works by restricting the IHT updates to an active set (a relatively small subset of variables) and occasionally augmenting the active set with variables that violate certain optimality conditions. By implementing this strategy, the iteration complexity of the algorithm can be reduced to O(nk)O(nk) in practice, resulting in an improvement, when kk is smaller than pp. The algorithm details can be found in Appendix A.3.

2.2 Determining a good stepsize

Choosing an appropriate stepsize τs\tau^{s} is crucial for fast convergence of the IHT algorithm. To ensure convergence to a stationary solution, a common choice is to use a constant stepsize of τs=1/L\tau^{s}=1/L (Bertsimas et al., 2016; Hazimeh & Mazumder, 2020), where LL is the Lipschitz constant of the gradient of the objective function. This approach, while reliable, can lead to conservative updates and slow convergence—refer to Appendix A.1 for details. An alternative method for determining the stepsize is to use a backtracking line search, as proposed in Beck & Teboulle (2009). The method involves starting with a relatively large estimate of the stepsize and iteratively shrinking the step size until a sufficient decrease of the objective function is observed. However, this approach requires multiple evaluations of the objective function, which can be computationally expensive.

Our novel scheme. We propose a novel line search method for determining the stepsize to improve the convergence speed of IHT. Specifically, we develop a method that (approximately) finds the stepsize that leads to the maximum decrease in the objective, i.e., we attempt to solve

For general objective functions, solving the line search problem (as in (10)) is challenging. However, in our problem, we observe and exploit an important structure: g(τs)g(\tau^{s}) is a piecewise quadratic function with respect to τs\tau^{s}. Thus, the optimal stepsize on each piece can be computed exactly, avoiding redundant computations (associated with finding a good stepsize) and resulting in more aggressive updates. In Appendix A.1, we present an algorithm that finds a good stepsize by exploiting this structure. Compared to standard line search, our algorithm is more efficient, as it requires fewer evaluations of the objective function and yields a stepsize that results in a steeper descent.

2.3 Additional techniques for scalability

While the IHT algorithm can be quite effective in identifying the appropriate support, its progress slows down considerably once the support is identified (Blumensath, 2012), resulting in slow convergence. We propose two techniques that refine the non-zero coefficients by solving sub-problems to speedup the overall optimization algorithm: (i) Coordinate Descent (CD, Bertsekas, 1997; Nesterov, 2012) that updates each nonzero coordinate (with others fixed) as per a cyclic rule; (ii) Back solve based on Woodbury formula (Max, 1950) that calculates the optimal solution exactly on a restricted set of size kk. We found both (i), (ii) to be important for improving the accuracy of the pruned network. Further details on the strategies (i), (ii) are in Appendix A.2 and A.4.

3 A multi-stage procedure

Our single-stage methods (cf Section 3.2) lead to high-quality solutions for problem (8). Compared to existing methods, for a given sparsity level, our algorithms deliver a better objective value for problem (8)—for eg, see Figure 2(b). However, we note that the final performance (e.g., accuracy) of the pruned network depends heavily on the quality of the local quadratic approximation. This is particularly true when targeting high levels of sparsity (i.e., zeroing out many weights), as the objective function used in (8) may not accurately approximate the true loss function L\mathcal{L}. To this end, we propose a multi-stage procedure named CHITA++ that improves the approximation quality by iteratively updating and solving local quadratic models. We use a scheduler to gradually increase the sparsity constraint and take a small step towards higher sparsity in each stage to ensure the validity of the local quadratic approximation. Our multi-stage procedure leverages the efficiency of the single-stage approaches and can lead to pruned networks with improved accuracy by utilizing more accurate approximations of the true loss function. For example, our experiments show that the multi-stage procedure can prune ResNet20 to 90% sparsity in just a few minutes and increases test accuracy from 15% to 79% compared to the single-stage method. Algorithm 1 presents more details on CHITA++.

Our proposed multi-stage method differs significantly from the gradual pruning approach described in Han et al. (2015). While both methods involve pruning steps, the gradual pruning approach also includes fine-tuning steps in which SGD is applied to further optimize the parameters for better results. However, these fine-tuning steps can be computationally expensive, usually taking days to run. In contrast, our proposed multi-stage method is a one-shot pruning method and only requires constructing and solving Problem (8) several times, resulting in an efficient and accurate solution. This solution can then be further fine-tuned using SGD or plugged into the gradual pruning framework, something we explore in Section 4.2.

Experimental Results

We compare our proposed framework with existing approaches, for both one-shot and gradual pruning.

We start by comparing the performance of our methods: CHITA (single-stage) and CHITA++ (multi-stage) with several existing state-of-the-art one-shot pruning techniques on various pre-trained networks. We use the same experimental setup as in recent work (Yu et al., 2022; Singh & Alistarh, 2020). The existing pruning methods we consider include MP (Magnitude Pruning, Mozer & Smolensky, 1989), WF (WoodFisher, Singh & Alistarh, 2020), CBS (Combinatorial Brain Surgeon, Yu et al., 2022) and M-FAC (Matrix-Free Approximate Curvature, Frantar et al., 2021). The pre-trained networks we use are MLPNet (30K parameters) trained on MNIST (LeCun et al., 1998), ResNet20 (He et al., 2016, 200k parameters) trained on CIFAR10 (Krizhevsky et al., 2009), and MobileNet (4.2M parameters) and ResNet50 (He et al., 2016, 22M parameters) trained on ImageNet (Deng et al., 2009). For further details on the choice of the Hessian approximation, we refer the reader to Appendix A.5. Detailed information on the experimental setup and reproducibility can be found in Appendix B.1.1.

Recent works that use the empirical Fisher information matrix for pruning purposes (Singh & Alistarh, 2020; Yu et al., 2022) show that using more samples for Hessian and gradient approximation results in better accuracy. Our experiments also support this conclusion. However, most prior approaches become computationally prohibitive as sample size nn increases. As an example, the Woodfisher and CBS algorithms require hours to prune a MobileNet when nn is set to 1000, and their processing time increases at least linearly with nn. In contrast, our method has been designed with efficiency in mind, and we have compared it to M-FAC, a well-optimized version of Woodfisher that is at least 1000 times faster. The results, as depicted in Figure 1, demonstrate a marked improvement in speed for our algorithm, with up to 20 times faster performance.

1.2 Accuracy of the pruned models

Comparison against state-of-the-art. Table 1 compares the test accuracy of MLPNet, ResNet20 and MobileNetV1 pruned to different sparsity levels. Our single-stage method achieves comparable results to other state-of-the-art approaches with much less time consumption. The multi-stage method (CHITA++) outperforms other methods by a large margin, especially with a high sparsity rate.

Sparsity schedule in multi-stage procedure. We study the effect of the sparsity schedule (i.e., choice of τ1≤τ2≤⋯≤τf=τ\tau_{1}\leq\tau_{2}\leq\cdots\leq\tau_{f}=\tau in Algorithm 1) on the performance of CHITA++. We compare test accuracy of three different schedules: (i) exponential mesh, (ii) linear mesh, and (iii) constant mesh. For these schedules, ff is set to be 1515. For the first two meshes, τ1\tau_{1} and τ15\tau_{15} are fixed as 0.20.2 and 0.90.9, respectively. As shown in Figure 3, the exponential mesh computes τ2,…,τ14\tau_{2},\dots,\tau_{14} by drawing an exponential function, while the linear mesh adopts linear interpolation (with τ1\tau_{1} and τ15\tau_{15} as endpoints) to determine τ2,…,τ14\tau_{2},\dots,\tau_{14} and the constant mesh has τ1=τ2=⋯=τ15\tau_{1}=\tau_{2}=\cdots=\tau_{15}.

Figure 4 plots the test accuracy of the three schedules over the number of stages. We observe that the linear mesh outperforms the exponential mesh in the first few iterations, but its performance drops dramatically in the last two iterations. The reason is that in high sparsity levels, even a slight increase in the sparsity rate leads to a large drop in accuracy. Taking small “stepsizes” in high sparsity levels allows the exponential mesh to fine-tune the weights in the last several stages and achieve good performance.

We perform additional ablation studies to further evaluate the performance of our method. These studies mainly focus on the effect of the ridge term (in Appendix B.2.1), and the effect of the first-order term (in Appendix B.2.2).

2 Performance on gradual pruning

To compare our one-shot pruning algorithms against more unstructured pruning methods, we plug CHITA into a gradual pruning procedure (Gale et al., 2019), following the approach in Singh & Alistarh (2020). Specifically, we alternate between pruning steps where a sparse weight is computed and fine-tuning steps on the current support via Stochastic Gradient Descent (SGD). To obtain consistent results, we start from the same pre-trained weights used in Kusupati et al. (2020), and re-train for 100 epochs using SGD during fine-tuning steps, similarly to Kusupati et al. (2020); Singh & Alistarh (2020). We compare our approach against Incremental (Zhu & Gupta, 2018), STR (Kusupati et al., 2020), Global Magnitude (Singh & Alistarh, 2020), WoodFisher (Singh & Alistarh, 2020), GMP (Gale et al., 2019), Variational Dropout (Molchanov et al., 2017), RIGL (Evci et al., 2020), SNFS (Dettmers & Zettlemoyer, 2020) and DNW (Wortsman et al., 2019). Further details on training procedure can be found in Appendix B.1.2.

MobileNetV1. We start by pruning MobileNetV1 (4.2M parameters). As Table 2 demonstrates, CHITA results in significantly more accurate pruned models than previous state-of-the-art approaches at sparsities 75% and 89%, with only 6% accuracy loss compared to 11.29%, the previous best result when pruning to a sparsity of 89%.

ResNet50. Similarly to MobileNetV1, CHITA improves test accuracy at sparsity levels 90%, 95%, and 98% compared to all other baselines, as Table 3 shows. This improvement becomes more noticeable as we increase the target sparsity, with CHITA producing a pruned model with 69.80% accuracy compared to 65.66%, the second-best performance, and previous state-of-the-art.

Conclusion

Acknowledgements

This research is supported in part by grants from the Office of Naval Research (N000142112841 and N000142212665), and Google. We thank Shibal Ibrahim for helpful discussions. We also thank Thiago Serra and Yu Xin for sharing with us code from their CBS paper (Yu et al., 2022).

References

Appendix A Algorithmic details

We propose a new line search strategy to efficiently determine an aggressive stepsize to address the issue of slow updates in the IHT algorithm. Note that the problem of finding the best stepsize can be written as the following one-dimensional problem

Since PkP_{k} is a piecewise function, g(τs)g(\tau^{s}) is a univariate piecewise quadratic function which is generally non-convex, as illustrated in Figure 5. Our key observation is that the first breaking point of g(τs)g(\tau^{s}) and the optimal stepsize on the first piece can be computed easily. More specifically, denote by τcs\tau_{c}^{s} the first breaking point of g(τs)g(\tau^{s}). Namely, τcs\tau_{c}^{s} is the largest value of τ′\tau^{\prime} such that the hard thresholding based on τs∈[0,τ′]\tau^{s}\in[0,\tau^{\prime}] does not change the support, i.e. Ik(wt)=Ik(wt−τs∇Q(wt)), ∀τs∈[0,τ′]\mathcal{I}_{k}(w^{t})=\mathcal{I}_{k}(w^{t}-\tau^{s}\nabla Q(w^{t})),~{}\forall\tau^{s}\in[0,\tau^{\prime}]. Let us denote S:=supp(w)\mathcal{S}:=\text{supp}(w). In the case where ∣S∣=k|\mathcal{S}|=k, τcs\tau_{c}^{s} can be computed in closed form using

As previously established, over the interval τs∈[0,τcs]\tau^{s}\in[0,\tau_{c}^{s}], the function g(τs)=Q(wt−τs∇Q(wt))g(\tau^{s})=Q\left(w^{t}-\tau^{s}\nabla Q(w^{t})\right) is a quadratic function. Let us denote by τms\tau_{m}^{s} the optimal value of τs\tau^{s} that minimizes g(τs)g(\tau^{s}) within the interval [0,τcs][0,\tau_{c}^{s}]. It is straightforward to see that τms\tau_{m}^{s} can be computed in closed form with the same computational cost as a single evaluation of the quadratic objective function.

If τms<τcs\tau_{m}^{s}<\tau_{c}^{s}, then the optimal value of τs\tau^{s} lies within the first quadratic piece. Practically, we have found that in this case τms\tau_{m}^{s} is often also the global minimum of g(τs)g(\tau^{s}). Therefore, we can confidently take the stepsize as τs=τms\tau^{s}=\tau_{m}^{s}. Otherwise if τms=τcs\tau_{m}^{s}=\tau_{c}^{s}, then we know that g(τs)g(\tau^{s}) is monotonically decreasing on the interval [0,τcs][0,\tau_{c}^{s}]. This implies that g(τs)g(\tau^{s}) would likely continue to decrease as τs\tau^{s} becomes larger than τcs\tau_{c}^{s}. As a result, we perform a line search by incrementally increasing the value of τs\tau^{s} by a factor of γ>1\gamma>1 starting from τcs\tau_{c}^{s} to approximate the stepsize that results in the steepest descent. The above procedure is summarized in Algorithm 2.

Our proposed scheme offers a significant improvement in efficiency compared to standard backtracking line search by eliminating redundant steps on the quadratic piece of g(τs)g(\tau^{s}) over [0,τcs][0,\tau_{c}^{s}]. Additionally, our method directly computes the optimal stepsize on the first piece, which in many cases, results in a greater reduction in the objective function when compared to the standard backtracking line search.

Finally, we note that during line search, it is always possible to find the piece of the quadratic function to which the current stepsize τs\tau^{s} belongs, say [τls,τus][\tau_{l}^{s},\tau_{u}^{s}], and calculate the optimal stepsize over that piece with small extra costs to further improve the line search. But we find it unnecessary in practice as the line search procedure usually terminates in a few steps.

A.2 Cyclic coordinate descent

Although IHT does well in identifying and updating the support, we observe that it makes slow progress in decreasing the objective in experiments. To address this issue, we use cyclic coordinate descent (CD, Bertsekas, 1997; Nesterov, 2012) with full minimization in every nonzero coordinate to refine the solution on the support. CD-type methods are widely used for solving huge-scale optimization problems in statistical learning, especially those problems with sparsity structure, due to their inexpensive iteration updates and capability of exploiting problem structure, such as Lasso (Friedman et al., 2010) and L0L2L_{0}L_{2}-penalized regression (Hazimeh & Mazumder, 2020).

Calculating CDUpdate(wt,i)\textsf{CDUpdate}(w^{t},i) requires the minimization of a univariate quadratic function with time cost O(n)O(n).

Cyclic CD enjoys a fast convergence rate (Bertsekas, 1997; Nesterov, 2012). However, the quality of the resulting solution is limited and depends heavily on the initial solution, as CD cannot modify the support of a solution. In practice, we adopt a hybrid updating rule that combines IHT and cyclic CD for better performance in terms of both quality and efficiency. In each iteration, we perform several rounds of IHT updates and then apply cyclic CD to refine the solution on the support. This approach is summarized in Algorithm 3.

A.3 Active set updates

The active set strategy is a popular approach that has been shown to be effective in reducing complexity in various contexts (Nocedal & Wright, 1999; Friedman et al., 2010; Hazimeh & Mazumder, 2020). In our problem setting, the active set strategy works by starting with an initial active set (of length equal to a multiple of the required number of nonzeros kk, e.g., 2k2k) that is selected based on the magnitude of the initial solution. In each iteration, we restrict the updates of Algorithm 3 to the current active set A\mathcal{A}. After convergence, we perform IHT updates on the full vector to find a better solution ww with supp⁡(w)⊈A\operatorname{supp}(w)\not\subseteq\mathcal{A}. The algorithm terminates if such ww does not exist; otherwise, we update A←A∪supp⁡(w)\mathcal{A}\leftarrow\mathcal{A}\cup\operatorname{supp}(w), and the process is repeated. Algorithm 4 gives a detailed illustration of the active set method, with Algorithm 3 as the inner solver (potentially the inner solver can be replaced with any other solver, such as Algorithm 5 in the next section). In our experiments, this strategy works well on medium-sized problems (p∼105)(p\sim 10^{5}) and sparse problems (k≪p)(k\ll p).

A.4 Backsolve via Woodbury formula

The backsolve method is stated in Algorithm 5.

We note that prior works (Singh & Alistarh, 2020; Yu et al., 2022; Hassibi & Stork, 1992) also use the formula, but they do not exploit the problem structure to reduce the runtime and memory consumption.

A.5 Stratified block-wise approximation

We describe in this subsection a block approximation strategy whereby we only consider limited-size blocks on the diagonal of the Hessian matrix and ignore off-diagonal parts. Given a disjoint partition {Bi}i=1c\{B_{i}\}_{i=1}^{c} of {1,2,…,p}\{1,2,\dots,p\} and assume blocks of size B1×B1,…,Bc×BcB_{1}\times B_{1},\dots,B_{c}\times B_{c} along the diagonal, problem (8) can then be decomposed into the following subproblems (1≤i≤c1\leq i\leq c)

where bi=ABiwˉBi−eb_{i}=A_{B_{i}}\bar{w}_{B_{i}}-e and ∑i=1cki=k\sum_{i=1}^{c}k_{i}=k determines the sparsity in each block. The difference in the selection of {ki}i=1c\{k_{i}\}_{i=1}^{c} will greatly affect the quality of the solution. We observe in experiments that the best selection strategy is to first apply magnitude pruning (or other efficient heuristics) to get a feasible solution ww, and then set ki=∣supp⁡(w)∩Bi∣, ∀1≤i≤ck_{i}=|\operatorname{supp}(w)\cap B_{i}|,\,\forall 1\leq i\leq c. Algorithm 6 states the block-wise approximation algorithm, with Algorithm 5 as the inner solver for each subproblem.

In our experiment, we adopt the same strategy to employ the block-wise approximation as in the prior work (Yu et al., 2022; Singh & Alistarh, 2020). We regard the set of variables that corresponds to a single layer in the network as a block and then subdivide these blocks uniformly such that the size of each block does not exceed a given parameter Bsize=104B_{size}=10^{4}.

We clarify that the introduction of block-wise approximation is for the sake of solution quality (accuracy of pruned network) rather than algorithmic efficiency. This differs from previous works (Singh & Alistarh, 2020; Yu et al., 2022). In fact, solving (17) for i=1,…,ci=1,\dots,c requires operations of the same order as solving (8) directly. On the other side, we observe in our experiments that adopting block-wise approximation will dramatically increase the network MobileNet’s accuracy (from 0.2% to near 30%, given a sparsity level of 0.8).

Appendix B Experiment details

All experiments were carried on a computing cluster. Experiments for MLPNet and ResNet20 were run on an Intel Xeon Platinum 8260 machine with 20 CPUs; experiments for MobileNetV1 and ResNet50 were run on an Intel Xeon Gold 6248 machine with 40 CPUs and one GPU.

We utilize the CHITA algorithm with active set strategy and coordinate descent as acceleration techniques, as outlined in Algorithm 4, to prune MLPNet and ResNet20 networks. Additionally, we use Algorithm 4 as the inner solver of our proposed multi-stage approach, CHITA++, for these networks. As to MobileNetV1 and ResNet50, we utilize CHITA-BSO with block approximation (Algorithm 6) for solving single-stage problems. We employ the exact block-wise approximation strategy as applied in previous work (Yu et al., 2022; Singh & Alistarh, 2020), see Section A.5 for details. We also use Algorithm 6 as the inner solver of our proposed multi-stage procedure CHITA++ for these networks. In our experiments, we set the number of stages in CHITA++ to 15 for MLPNet and ResNet20 and 100 for MobileNetV1. CHITA++ results on MobileNetV1 are averaged over 4 runs.

For each network and each sparsity level, we run our proposed methods CHITA (single-stage) and CHITA++ (multi-stage) with ridge value λ\lambda ranging from [10−5,103][10^{-5},10^{3}] and the number of IHT iterations (if Algorithm 4 is applied) ranging from $$. In single-stage settings, we consider solving problem (8) with/without the first-order term. We report in Table 1 the best model accuracy over all possible hyper-parameter combinations.

To obtain consistent results, we run CHITA and M-FAC with the same set of hyperparameters (λ=10−5,n=500,Bsize=104\lambda=10^{-5},n=500,B_{size}=10^{4}) and on the same training samples for Hessian and gradient approximation. We performed a sensitivity analysis with different block sizes BsizeB_{size} and found similar results — suggesting that the results are robust to the choice of BsizeB_{size}.

B.1.2 Gradual pruning

All experiments were carried on a computing cluster. Experiments for MobileNetV1 were run on an Intel Xeon Platinum 6248 machine with 30 CPUs and 2 GPUs; experiments for ResNet50 were run on five Intel Xeon Platinum 6248 machines with 200 CPUs and 10 GPUs.

In all our gradual pruning experiments, we begin by pruning the networks to a sparsity level of 50% and proceed with six additional pruning steps to reach the target sparsity. We follow the polynomial schedule introduced by Zhu & Gupta (2018) as the pruning schedule and use the CHITA-BSO algorithm with block approximation (Algorithm 6) as the pruning method. The block size is set to Bsize=2000B_{size}=2000 for MobileNetV1 and Bsize=500B_{size}=500 for ResNet50.

We incorporate SGD with a momentum of 0.90.9 for 12 epochs between pruning steps. Once the networks have been pruned to the target sparsity, we continue to fine-tune the networks for an additional 28 epochs using SGD with a momentum of 0.90.9 (total of 100100 epochs). We utilize distributed training and set the batch size to 256256 per GPU during the SGD training process.

We implement a cosine-based learning rate schedule similar to the one used in the STR method (Kusupati et al., 2020). Specifically, the learning rate for each epoch ee between two pruning steps that occur at epochs e1e_{1} and e2e_{2} is defined as:

Figure 6 illustrates how such a learning rate schedule decays between pruning steps.

B.2 Implementation details and ablation studies

In this section, we study the effect of the ridge term on the performance of our algorithm, specifically focusing on the test accuracy over the course of the algorithm. As depicted in Figure 7(a), when no ridge term is applied, the test accuracy increases initially but then experiences a sharp decline as the algorithm progresses. The underlying cause is revealed in Figure 7(b), which illustrates that without the ridge term, the distance between the original weight wˉ\bar{w} and the pruned weight ww keeps increasing as the algorithm progresses. As this distance increases, the local quadratic model used in (8) becomes less accurate, leading to poor test performance.

One solution to this problem would be to employ early stopping to prevent the distance from growing too large. However, determining the optimal stopping point can be challenging. Practically, we instead add the ridge term nλ2∥w−wˉ∥2\frac{n\lambda}{2}\|w-\bar{w}\|^{2} to the objective function, effectively regularizing the model and maintaining its accuracy. As shown in Figure 7(c), utilizing a well-tuned ridge term results in an increase of approximately 3% on MLPNet.

B.2.2 Effect of the first-order term

In scale-independent applications, e.g., minimizing L(w)≈L(wˉ)+12(w−wˉ)⊤H(w−wˉ)\mathcal{L}(w)\approx\mathcal{L}(\bar{w})+\frac{1}{2}(w-\bar{w})^{\top}H(w-\bar{w}) as considered in Singh & Alistarh (2020) and Yu et al. (2022), the empirical Fisher matrix HH still effectively approximates the true Hessian. However, this approximation is no longer accurate in our framework, which includes a first-order term. This is supported by the results shown in Figure 8(a), where our framework with a correctly scaled term (α=1/m\alpha=1/m) demonstrates significantly improved performance compared to one without a scaling factor (α=1\alpha=1), especially when the fisher batch size mm is much greater than one.

To address this issue, we propose a local quadratic approximation with a scaled first-order term that reads

However, the computation cost of Trace(∇2L(wˉ))\text{Trace}(\nabla^{2}\mathcal{L}(\bar{w})) is not negligible, even using accelerated methods as proposed in Yao et al. (2020). Through experimentation, as shown in Figure 8(b), we have discovered that the estimated value of α\alpha as given by (20) is relatively close to 1/m1/m. Therefore, we have chosen to use 1/m1/m as a heuristic scaling factor in our experiments, as it provides a good approximation while reducing the computational cost.

In Figure 8(c), we further illustrate the benefits of using large mini-batches and a scaled first-order term. As the fisher batch size mm increases, we can construct more precise local quadratic approximations through better estimation of HH and gg, resulting in improved test accuracy. Additionally, when mm is greater than 1, using a correctly scaled first-order term provides an additional performance boost.