The Optimal BERT Surgeon: Scalable and Accurate Second-Order Pruning for Large Language Models

Eldar Kurtic, Daniel Campos, Tuan Nguyen, Elias Frantar, Mark Kurtz, Benjamin Fineran, Michael Goin, Dan Alistarh

Introduction

Pre-trained Transformer models Vaswani et al. (2017); Devlin et al. (2019) provide robust language representations which can be specialized on various tasks. Given their massive growth Radford et al. (2019); Smith et al. (2022), techniques for reducing their computational overheads have become popular. One classic technique is Knowledge Distillation (KD) Hinton et al. (2015), which transfers knowledge from a larger teacher to a smaller student model. Other work has leveraged lower-precision representations to produce quantized models. An orthogonal approach, which is our primary focus, has been to apply unstructured and block pruning, i.e. removing individual weights, to produce compressed but accurate language models. Figure 1 provides a comparative overview of state-of-the-art results for unstructured pruning.

In this paper, we introduce a method for improved unstructured and semi-structured (block) pruning, by leveraging the second-order approach pioneered by the Optimal Brain Surgeon framework LeCun et al. (1989); Hassibi and Stork (1992), which we scale for the first time to LLMs. Further, we put our results in the context of a compound compression approach, which combines several compression techniques to obtain sparse models which we execute on a sparsity-aware CPU-based runtime NeuralMagic (2021), showing order-of-magnitude speedups at low accuracy loss.

In summary, our contributions are as follows:

We perform a thorough exploration of weight pruning approaches applied to LLMs, including lottery-ticket, movement pruning, magnitude and second-order pruning.

We introduce a general second-order pruning method called Optimal BERT Surgeon (oBERT), which supports unstructured and block pruning, and is the first second-order method to be both highly-accurate and scalable to the dimensionality of BERT models.

We illustrate the benefits of oBERT by significantly improving upon existing state-of-the-art pruning methods, in both stages of language tasks: pre-training and fine-tuning. For illustration, when pruning BERTBASE \textrm{BERT}_{\textrm{BASE}}\,, oBERT outperforms Movement Pruning (MvP), the most accurate prior approach, by more than 2% absolute F1 score at the same sparsity, and can match the accuracy of MvP models with 3x fewer parameters.

We investigate the applicability of this pruning method in a framework which compounds popular compression approaches for LLMs, i.e. applying pruning in combination with layer dropping and/or quantization. In this context, we show that our resulting sparse models provide order-of-magnitude improvements compared to other compound compressed models, and that they can be easily deployed for CPU inference.

Background and Related Work

Transformer LLMs are usually built using multiple transformer layers with self-attention Vaswani et al. (2017). Each transformer has a variation of two sub-components: multi head attention (MHA) and fully connected feed forward network (FFN). Given the massive size of well-performing models, there has been growing interest in LLM compression. They have been shown to be fragile as minor perturbations can lead to model collapse Kovaleva et al. (2021). Pruning schemes are motivated by weight saliency metrics which represent the loss in accuracy due to pruning. It is common to prune in iterative steps, each of which removes weights until a desired sparsity level is reached. Now, we briefly overview existing approaches.

Structured pruning for LLMs focuses on reducing the number of layers and/or attention heads, and requires structural understanding of the model. Michel et al. (2019) and Voita et al. (2019) demonstrated that for some tasks nearly 40% of attention heads can be removed without major impact on accuracy. Other work has focused on removing layers Sridhar and Sarah (2020), and on the order in which they are removed Sajjad et al. (2020). In some of our experiments, we apply standard “direct” layer dropping in conjunction with pruning.

Semi-structured pruning is an intermediate approach, by which smaller groups, e.g. rectangular sets of weights Lagunas et al. (2021), are set to zero. This approach has recently gained in popularity thanks to efficient computational support. We extend the second-order pruning formulation to such groupings, and show results for a specific grouping supported by a CPU-inference engine.

Unstructured pruning removes individual weights by setting them to zero. Gradual Magnitude Pruning (GMP) is a classic approach, which makes use of weight magnitudes as a saliency metric for pruning Han et al. (2015); Gale et al. (2019).

First-order pruning methods use a gradient based formulation of the saliency metric. A popular method is Movement Pruning (MvP) Sanh et al. (2020), specifically designed for pruning in the fine-tuning stage. Intuitively, it removes weights that are moving towards zero. The resulting models were the first to achieve high sparsity with tolerable accuracy loss. Methods such as PLATON Zhang et al. (2022) attempt to capture the uncertainty of weights importance scores by upper confidence bound estimation. Prior to our work, MvP and PLATON approaches set state-of-the-art results for unstructured pruning.

Second-order pruning methods LeCun et al. (1989); Hassibi and Stork (1992); Singh and Alistarh (2020); Frantar et al. (2021) were developed in the context of image classification, and leverage complex approximations of the loss curvature. However, second-order pruning methods require an approximation of the inverse Hessian, which is expensive to store and compute with for LLM parameter counts. The approach we propose is similar to WoodFisher/M-FAC methods Singh and Alistarh (2020); Frantar et al. (2021), but is the first to work accurately at LLM scale. Specifically, the WoodFisher approach is infeasible at BERT scale, as it requires storing gradients for inverse Fisher calculation in memory at the point of pruning. The M-FAC approach scales, but we show that its parametrization yields worse pruning results (Appendix Figure 3). This is because M-FAC performs full-matrix (non-blocked) inversion by default, which is inherently noisy. In addition, we extend the theoretical OBS approach to semi-structured (block) compression. We also show that our method can be applied during LLM pre-training and fine-tuning, yielding state-of-the-art results in both regimes.

Knowledge Distillation Hinton et al. (2015) trains a smaller student model against outputs of a larger teacher model by adding a loss component which minimizes the KL-divergence between the two output distributions, which is the approach we adopt in our setup too. A hardness parameter is used to control the mixture of regular and distillation loss, and a temperature parameter to control softness of the distribution. Contrary to this, approaches like DistilBERT Sanh et al. (2019), TinyBERT Jiao et al. (2020), MobileBERT Sun et al. (2020a), and MiniLM Wang et al. (2020) utilize more complex distillation schemes, based on transferring knowledge from intermediate model’s representations. Our sparse models provide order-of-magnitude improvements upon some of these methods.

Quantization represents weights and activations in lower precision Courbariaux et al. (2016), and was used to obtain models such as Q8BERT Zafrir et al. (2019) and TernaryBERT Zhang et al. (2020).

Shen et al. (2020) uses information about the Hessian spectrum to choose quantization bit-widths, whereas Yu et al. (2022) uses an approximation of the Hessian trace for structured pruning. These Hessian-based approaches are different from the one we propose, as we use completely different inverse-Hessian approximations to guide pruning decisions. The focus of our work is on weight pruning, and on computational speedups achievable on commodity CPUs. As such, the methods we investigate are orthogonal to quantization. Moreover, it is impossible to directly compare to low-bitwidth quantized models as most inference frameworks do not support such custom formats. Therefore, we will only make use of the standard Quantization-Aware Training (QAT) to 8-bit weights, which is well-supported on Intel CPUs, and showcase the resulting speedups in conjunction with layer dropping and weight pruning.

Downstream compression methods attempt to compress directly while fine-tuning on a specific task. MvP method is specially designed for this setup. Upstream compression methods compress during the pre-training phase, reducing the need for task-specific pruning. Chen et al. (2020) examined the “Lottery Ticket” strategies Frankle and Carbin (2018) which, as we illustrate later, incur huge accuracy loss even at moderate sparsities. Recent work “Prune Once for All” (Prune OFA) by Zafrir et al. (2021) showed that well-tuned magnitude pruning can be competitive with downstream methods like MvP.

We first examine the performance of prior pruning methods, notably MvP, Prune OFA, and Lottery Tickets, relative to the new second-order oBERT method. The approach we propose consistently improves upon all these prior methods, both in pre-training (upstream) and fine-tuning (downstream) stages, and can be compounded with other compression techniques to obtain models that are smaller, faster and more accurate than models like DistilBERT, TinyBERT, and block MvP.

Additional approaches for efficient inference of LLMs exist, like token-pruning and early-exiting. These approaches are orthogonal to ours; therefore we discuss them in Appendix A.2.

The Optimal BERT Surgeon (oBERT)

Given that w∗\mathbf{w}^{*} is well-optimized, it is reasonable in practice to assume that ∇L(w∗)≈0\nabla\mathcal{L}(\mathbf{w}^{*})\approx\mathbf{0}. Then, the change in loss incurred by pruning a subset of weights can be expressed as

where δL(δw)≔L(wM)−L(w∗)\delta\mathcal{L}(\delta\mathbf{w})\coloneqq\mathcal{L}(\mathbf{w}_{M})-\mathcal{L}(\mathbf{w}^{*}) and δw≔wM−w∗\delta\mathbf{w}\coloneqq\mathbf{w}_{M}-\mathbf{w}^{*}. A popular way of approximating the Hessian at w∗\mathbf{w}^{*} is via a dampened empirical Fisher information matrix Hassibi and Stork (1992):

Returning to our pruning problem, assume we wish to identify a block of weights QQ of a given shape whose removal by zero-masking would incur minimum increase in loss. This leads to the following constrained optimization problem:

which prunes a set of weights QQ and updates the remaining weights to preserve the loss. Now, the corresponding loss increase incurred by the optimal weight update δw∗\delta\mathbf{w}^{*} can be expressed as the saliency score of weights QQ, which we denote by:

We use this saliency/importance score to rank groups of weights for pruning. As a sanity check, if we prune a single weight wjw_{j} at a time, our derivations will yield the standard formulas of Hassibi and Stork (1992). The full version of Singh and Alistarh (2020) provided a slightly less general derivation for the blocked case, under additional assumptions.

2 An Efficient Implementation

Assume a gradual pruning setup, in which at each pruning step we wish to prune a model to a target sparsity s∈(0,1]s\in(0,1], effectively zeroing out s×ds\times d weights, in groups of size ∣Q∣|Q|. Typically s×d≫∣Q∣s\times d\gg|Q|, meaning that we want to remove multiple groups at the same time. Finding the optimal set of s×d∣Q∣\frac{s\times d}{|Q|} groups is an intractable combinatorial problem, due to all possible correlations between them, given by the binomial coefficient (nk)\binom{n}{k}, where n=d∣Q∣n=\frac{d}{|Q|} and k=s×d∣Q∣k=\frac{s\times d}{|Q|}. This problem can be alleviated by ignoring correlations between different groups of weights QQ, and solving only for correlations between the weights within the same group. In practice, this boils down to evaluating the saliency score ρQ\rho_{\textrm{Q}} for each group QQ, and pruning the s×d∣Q∣\frac{s\times d}{|Q|} groups with the lowest score. As pruning many weights in the same step can make the Taylor approximation of the loss function less accurate, one can consider pruning with multiple smaller sub-steps with recomputations of the Hessian approximation in between (without intermediate fine-tuning). While this can further improve the quality of the pruning step Frantar et al. (2021), we do not implement this additional optimization since the competing methods do not utilize recomputations.

2.2 Inverse empirical Fisher computation

The key space and time complexity cost of the above procedure is computing products with the inverse empirical Fisher. A direct approach would be to perform a block-wise diagonal approximation of this matrix (which we detail next), and perform direct block inversion. However, we found experimentally that this approach is too expensive in terms of time, and quite numerically-sensitive. As an alternative, we rely on the fact that the matrix we wish to invert is a sum of rank-1 matrices, and employ the Woodbury/Sherman-Morrison (WSM) inversion formula. Specifically, given a sum (A+uv⊤)(\mathbf{A}+\mathbf{u}\mathbf{v}^{\top}) of an invertible matrix A\mathbf{A} and an outer product of vectors u\mathbf{u} and v\mathbf{v} with compatible dimensions, the inverse (A+uv⊤)−1(\mathbf{A}+\mathbf{u}\mathbf{v}^{\top})^{-1} can be exactly calculated as A−1−A−1uv⊤A−11+v⊤A−1u\mathbf{A}^{-1}-\frac{\mathbf{A}^{-1}\mathbf{u}\mathbf{v}^{\top}\mathbf{A}^{-1}}{1+\mathbf{v}^{\top}\mathbf{A}^{-1}\mathbf{u}}. Placing the expression of the empirical Fisher in the WSM formula, we obtain the following recursive formulation, where mm is the number of gradients employed in the approximation:

Unrolling the recursion with F^0−1(w)=1λId\widehat{\mathbf{F}}^{-1}_{0}(\mathbf{w})=\frac{1}{\lambda}\mathbf{I}_{d}, we can obtain an iterative formula to exactly calculate the inverse of the empirical Fisher matrix as

The iterative formulation enjoys a number of computational advantages over the direct implementation. The most notable ones are 1) avoiding explicit calls to the expensive and dampening-sensitive matrix inversions, and 2) allowing successive updates of the inverse as new gradients are computed, never needing to store all mm gradients of size dd and thus significantly reducing memory requirements.

3 Memory and run-time complexity

Another alternative we investigated was the matrix-free approach of Frantar et al. (2021), which does not require a block-wise approximation and has complexity Θ(dm)\Theta(dm). However, our investigation showed that this approach required high values of mm to be accurate (Appendix Figure 3), which leads to excessive memory cost in the case of BERT models.

4 Efficient and scalable implementation

Experimental Validation

To ease reproducibility, we conduct our experiments in modified versions of the popular open-source libraries: Transformers Wolf et al. (2020), and SparseML Kurtz et al. (2020). All of our experiments are using publicly available datasets via Lhoest et al. (2021) and focus on the BERTBASE \textrm{BERT}_{\textrm{BASE}}\,model Devlin et al. (2019), one of the most commonly used LLMs, composed of 12 transformer layers with 110M parameters. Following community standards, we prune encoder’s weights (85M) and report sparsities relative to this number. All of our models, compression recipes and the full implementation will be made public.

We first revisit the accuracy-compression trade-off for pruning on downstream tasks.

Goals and setup. We compare existing approaches, notably Movement Pruning (MvP) Sanh et al. (2020) and Lottery Ticket (LT-BERT) Chen et al. (2020), against the gradual unstructured oBERT method, introduced in Section 3. Our experiments evaluate performance on a variety of downstream (English) tasks commonly used to evaluate model compression: question answering SQuAD v1.1 Rajpurkar et al. (2016), sentence classification Quora Duplicate Query Dataset QQP Shankar et al. (2017), and natural language inference MNLI Williams et al. (2018).

Comparison with MvP. For a fair comparison with MvP, we consider the 10-epoch gradual pruning setup used to obtain the best results by Sanh et al. (2020). Specifically, we start from the BERTBASE \textrm{BERT}_{\textrm{BASE}}\,model and perform 2 epochs of fine-tuning, followed by 6 epochs of pruning, and 2 further epochs of fine-tuning of the compressed model. We impose a global sparsity distribution over all layers, prune with oBERT two times per epoch, and use KD from the fine-tuned BERTBASE \textrm{BERT}_{\textrm{BASE}}\,teacher. For oBERT pruning we use m=1024m=1024 gradients, block size B=50B=50, and dampening λ=10−7\lambda=10^{-7} to approximate the inverse Hessian matrix. In all of our runs, the first pruning step prunes 70% of weights and then follows the cubic interpolation Zhu and Gupta (2018) to the target sparsity. This large first pruning step gives more time to recover from the later pruning steps, which impose higher sparsities. All hyper-parameters are described in detail in Appendix A.5, and the results are given in Table 1 (in the 10 Epochs section).

We observe that Optimal BERT Surgeon outperforms Movement Pruning by a significant margin, more than 2 points of F1 score at the same sparsity. Remarkably, the model pruned with oBERT to 97% sparsity has similar accuracy to MvP-pruned model at 90% sparsity, which has roughly 3x more weights. This reinforces the effectiveness of second-order information for pruning.

Extended pruning and fine-tuning. Next, we examine effects of extending the gradual schedule to 30 epochs, matching the setup used for LT-BERT Chen et al. (2020). The only difference compared to our 10 epoch setup is that we now prune with oBERT every four epochs, and rewind learning rate after each pruning step. The extended setup leaves more time to recover from pruning, which reflects in the improved results in Table 1 (30 Epochs section). We report the mean over three runs. For additional evaluation metrics and standard deviations please see Tables 12 and 15 in the Appendix. The results show a clear accuracy difference between oBERT and LT-BERT, especially at high sparsities. This difference is justified since the LT based approach attempts to mainly transfer network connectivity, whereas the oBERT can also benefit from the weight values. Finally, we examined the impact of extended setup with Soft MvP on SQuAD, targeting 90% sparsity (not shown in the Table), leading to an (F1, EM) combination of (87.42,79.83)(87.42,79.83) for MvP. The F1 gap in favor of oBERT is lower than at 10 epochs, suggesting that extended finetuning helps all methods; yet, it is far from negligible.

2 Upstream Unstructured Pruning

An appealing alternative to downstream pruning is to compress models upstream, on the semi-supervised pre-training task Zafrir et al. (2021). Given the upstream pruned model, computational requirements for obtaining downstream fine-tuned models are significantly reduced, as only fine-tuning of the remaining weights is necessary.

Goals and setup. To compare with existing approaches, notably Prune OFA Zafrir et al. (2021) and LT-BERT Chen et al. (2020), we gradually prune with oBERT directly at upstream datasets, BookCorpus and English Wikipedia, and then fine-tune the remaining unpruned weights on the subset of GLUE tasks.

Teacher preparation. Following Liu et al. (2019), we start with the HuggingFace BERTBASE \textrm{BERT}_{\textrm{BASE}}\,uncased model, and fine-tune it for additional 10 epochs only on the masked language modeling task.

Pruning at upstream. Once the distillation teacher is trained, we gradually prune and fine-tune the BERTBASE \textrm{BERT}_{\textrm{BASE}}\,model for 3 epochs, using KD from the dense teacher. We prune four times per epoch, and rewind learning rate to the initial value after each pruning step. Hyper-parameters for oBERT are the same as for downstream pruning in 4.1; a full description can be found in Appendix A.6.

Sparse-transfer to downstream. To evaluate the resulting upstream-pruned models, we finetune the unpruned weights on downstream tasks with KD from the fine-tuned BERTBASE \textrm{BERT}_{\textrm{BASE}}\,model. For a fair comparison with Prune OFA, we fine-tune for 8 epochs. The results in Table 2 show that sparse models produced by oBERT outperform state-of-the-art methods by significant margins. We report the mean over four runs. For additional evaluation metrics and standard deviations please see Appendix Tables 13 and 16. It is worth emphasizing that in contrast to Prune OFA, which performed extensive hyper-parameter tuning for sparse-transfer, our recipe is simple and general across downstream tasks: 8 epochs of fine-tuning with linearly decaying learning rate. This suggests that sparse pre-trained models found by oBERT constitute a strong starting point for sparse transfer learning, which can be further improved by task-specific hyper-parameter tuning.

3 Compound Compression for CPUs

To probe the potential practical impact of our approach, we specialize the technique for deployment on CPUs, corresponding to “edge” deployments. Specifically, we tailor our sparse models to the DeepSparse NeuralMagic (2021) sparsity-aware runtime, by compounding unstructured pruning with additional compression techniques.

Direct layer dropping. The competitive results obtained at high sparsities in sections 4.1 and 4.2 suggest that BERTBASE \textrm{BERT}_{\textrm{BASE}}\,may be overparameterized for downstream tasks. To improve compression ratio and inference speed, we apply “direct” layer dropping: we initially drop all but 3 or 6 of the BERT’s 12 layers. We drop layers from our upstream teacher, and, following Turc et al. (2019), fine-tune them with KD in the same setup used to prepare the upstream teacher. These 3 and 6 layer models are used as starting points for downstream pruning. More sophisticated layer dropping techniques Fan et al. (2019), could bring further accuracy gains; we leave this for future work.

Block pruning and QAT. High-performance inference usually benefits more from (semi) structured sparsity patterns than from the unstructured ones. Hence, we employ the generalized oBERT formulation introduced in the section 3 and prune weights in the 4-block pattern, meaning that contiguous blocks of 4 weights are either set to zero or kept dense. Both pruning types, unstructured and 4-block, can be leveraged for computational speedups with the DeepSparse runtime, but 4-block pruning coupled with INT8 quantization can provide further performance gains. For quantization, we apply standard quantization-aware training (QAT) Jacob et al. (2018) on top of the 4-block models (see Appendix A.7 for a full description).

Compounding for deployment. To determine the impact of different compression schemes, we investigate unstructured and 4-block pruning of the 3, 6, and 12-layer models. For all runs, we use the same set of hyper-parameters from the extended pruning and fine-tuning setup in Section 4.1. The results are given in Table 3, where we also report accuracy of the corresponding dense models (0% sparsity) in the same setup. For additional evaluation metrics, please see Table 14. The results indicate that compression methods can be combined without model collapse, although the accuracy drops do compound. The fact that the layer-dropped models are also highly compressible suggests that structured and fine-grained (unstructured) compression are complementary. We find it remarkable that our 6-layer unstructured oBERT-pruned model is competitive with the 12-layer MvP-pruned model when both are pruned to 90% sparsity.

Practical trade-offs. We now benchmark these models in end-to-end fashion, both in terms of model size and inference speed. For model size, we report size of the checkpoint in MB after standard gzip compression. For inference speed, we report number of items per second (throughput) on the well-established SQuAD v1.1 CPU-inference benchmark with a sequence length of 128 and a batch size of 32. Figure 2 depicts relative accuracy versus magnitude of improvement in speed and model size. As baseline for full recovery, we follow the community-standard e.g. Sanh et al. (2020), and adopt the dense BERTBASE \textrm{BERT}_{\textrm{BASE}}\,model with 88.54 F1 score. The baseline for inference speed is dense BERTBASE \textrm{BERT}_{\textrm{BASE}}\,inference with DeepSparse, which matches the industry-standard ONNX Runtime inference engine. Results suggest a roughly-linear trade-off between compression and accuracy loss, with a compression jump around 1% accuracy drop, due to quantization being applied. Specifically, we observe 8.4x higher inference speedup at < 1% accuracy drop, 10x speedup at < 2% drop, 15x speedup at < 3% drop, and 29x speedup at < 7.5% accuracy drop. This shows how compound compression can optimize LLMs to various latencies. See Appendix Table 17 for full results.

4 Pruning for GPU speedups (N:M sparsity)

Even though our previous results targeted CPUs for deployment, we now show that our pruning approach can also be relevant to GPUs. We apply the semi-structured variant of oBERT to impose the 2-out-of-4 sparsity pattern, which is supported on NVIDIA Ampere GPUs (Mishra et al., 2021). More specifically, we prune in one-shot, and compare against the magnitude pruning baseline in Table 4. All other methods require full fine-tuning, and thus don’t support the one-shot setup. oBERT significantly outperforms magnitude pruning, and with only 1-epoch of fine-tuning it is able to fully recover dense accuracy with (F1, EM) = (88.58, 81.16). With this sparsity pattern, the pruned model achieves 1.85x speedup on Ampere devices.

Discussion

Comparison with concurrent work. Concurrent work introduced PLATON Zhang et al. (2022), which addresses unstructured pruning of BERT models via estimates of confidence bounds. It does not make use of KD, so for a fair comparison we rerun our experiments without KD as well. Contrary to PLATON, which reports best results after an extensive hyper-parameter search for each task independently, we apply our sparse-transfer setup with the upstream pruned model and only sweep for the number of epochs ∈\in. We employ early stopping to prevent overfitting on smaller GLUE tasks. As can be seen from Table 5, oBERT outperforms PLATON across all tasks.

Broader comparison. We now contrast our compound-compressed BERTBASE \textrm{BERT}_{\textrm{BASE}}\,models relative to alternative compression techniques. We compare against DistilBERT Sanh et al. (2019), TinyBERT Jiao et al. (2020), and Block Pruning For Faster Transformers (Hybrid Filled MvP) Lagunas et al. (2021). DistilBERT leverages KD during pre-training and fine-tuning to obtain a 6-layer model fine-tuned for a specific downstream task. TinyBERT makes use of a specialized Transformer-KD scheme to distill knowledge and intermediate representations at both stages, pre-training and fine-tuning on a specific task. In contrast, we use a simpler approach and employ KD from teacher’s outputs only. Hybrid Filled MvP Lagunas et al. (2021) employs semi-structured pruning and weight reintroduction. The comparison is given in Table 6, where we report the number of unpruned encoder weights as size, compression ratio and inference speedup relative to the dense BERTBASE \textrm{BERT}_{\textrm{BASE}}\,in the same inference environment, and F1 score on the dev-set of the SQuAD v1.1 dataset. The results suggest that our compressed models improve upon the current state-of-the-art techniques, setting new very competitive baselines with respect to all metrics: accuracy, model size, and inference speed.

BERTLARGE \textrm{BERT}_{\textrm{LARGE}}\,results. Most of our results presented in Section 4 targeted the widely-adopted BERTBASE \textrm{BERT}_{\textrm{BASE}}\,model. This gave us an opportunity for a fair comparison against many different methods. To verify that our approach does not pertain only to the BERTBASE \textrm{BERT}_{\textrm{BASE}}\,model, in Table 7 we present downstream pruning results on the three times larger BERTLARGE \textrm{BERT}_{\textrm{LARGE}}\,model and the SQuADv1.1 task. As can be seen from the Table, even the model pruned with oBERT at double the sparsity (95%) outperforms Prune OFA (90%).

MLPerf Inference Benchmark. Motivated by our state-of-the-art results across-the-board, we apply our full compound compression pipeline to compress BERTLARGE \textrm{BERT}_{\textrm{LARGE}}\,and MobileBERT Sun et al. (2020b) models in the context of the industrial MLPerf Inference Benchmarkhttps://mlcommons.org/en/. In brief, we were able to achieve order-of-magnitude improvements in terms of model size and inference speedups, while maintaining >99% of the dense BERTLARGE \textrm{BERT}_{\textrm{LARGE}}\,accuracy. For details please see Appendix A.1, as well as our open-source submission.

Broader Impact

Our work is part of the general trend of producing inference efficient models which approximate performance of their larger bases. By and large, this work should help increase model efficiency, thereby reducing computational and ultimately monetary cost of executing such models. Moreover, it could allow models to be used by those who do not have access to expensive specialized computing clusters: for instance, our main speedup results are aimed at widely-available CPUs.

Limitations

As any academic study, our work is not without its limitations. We split their discussion into limitations that are inherent to our method, and limitations of our present study; the latter can be overcome by extensions of our work. In the first category, we begin by highlighting the fact that our second-order method relies on approximations, which are inherent in order to scale such methods to BERT scale. Prior studies, e.g. Singh and Alistarh (2020) have performed careful examinations of the validity of these approximations in the context of CNN models. The strength of our empirical results can be seen as indirect evidence that these approximations apply to BERT models as well. A second, technical, limitation is the fact that our method requires non-trivial additional storage cost; while we have shown that our experiments can be executed on a single commodity GPU (NVIDIA RTX 3090), this limits the range of devices on which the technique may be applied. However, we provide an efficient and easy way to scale our approach with more GPUs, which is automatically utilized in a multi-GPU environment.

Another limitation which we aim to remove in future work is the focus on relatively fine-grained sparsity types, such as unstructured and semi-structured pruning.

References

Appendix A Appendix

Following the MLPerf benchmark guidelines on producing compressed and fast models while maintaining >99% of the BERTLARGE \textrm{BERT}_{\textrm{LARGE}}\,F1 score on the SQuADv1.1 task, we explore two directions. In the first one, dubbed oBERT-Large, we compound compress the BERTLARGE \textrm{BERT}_{\textrm{LARGE}}\,model without any changes to its architecture. Therefore, we apply 4-block downstream pruning to 95% sparsity followed by the quantization aware training (QAT). In the second direction we focus on recovering the BERTLARGE \textrm{BERT}_{\textrm{LARGE}}\,accuracy by compressing an already compact MobileBERT model, dubbed oBERT-MobileBERT. More specifically, we apply direct layer dropping, leaving only 14 transformer layers out of the original 24, followed by the 4-block pruning to 50% sparsity and quantization aware training. We present results in Table 8, where models were evaluated with the DeepSparse inference engine, using a server with two Intel(R) Xeon(R) Platinum 8380 (IceLake) CPUs with 40 cores each, batch-size 128 and sequence length 384. For more details please see our official submission at https://github.com/neuralmagic/mlperf_inference_results_v2.1/tree/master/open/NeuralMagic.

A.2 Additional comparisons

Here we reflect upon some other methods focused on efficient inference for LLMs, which are orthogonal to weight pruning. For example, Learned Token Pruning Kim et al. (2022) tries to adaptively remove unimportant tokens in input sequences and provides 2x higher throughput at < 1% accuracy drop; at the same accuracy drop, our compressed model is able to achieve 8.4x higher throughput. DeeBERT Xin et al. (2020) and FastBERT Weijie et al. (2020) apply an early-exit technique for inference speedup. The latter achieves 2-3x faster inference without performance degradation. However, the method only applies to batch size one. Nevertheless, in terms of direct comparison, our compressed models are able to achieve 4x faster inference on CPUs without accuracy degradation. Overall, we emphasize the fact that these methods are complementary to our compression techniques, so it would be interesting to investigate computational gains by combining such methods.

A.3 Computational costs

In practice, for the 12-layer BERTBASE \textrm{BERT}_{\textrm{BASE}}\,model with d=85Md=85M encoder weights and block size B=50B=50, the O(Bd)\mathcal{O}(Bd) memory requirement translates to approximately 17GB, which can be easily kept on the 24GB RTX 3090 card. While this amount of memory is available on high-performance GPUs, it is also straightforward to split the NB×B×BN_{B}\times B\times B tensor along the batch-dimension NBN_{B} and utilize additional GPUs or even memory swapping with CPU. Our implementation updates the inverse Hessian approximation in negligible time, and can run asynchronously while the next gradient is being fetched. Computing saliency scores and optimal weight updates takes only a few seconds.

A.4 Optimal BERT Surgeon (oBERT) hyper-parameters

Hyper-parameters. The oBERT pruning method has three tunable hyper-parameters: number of gradients (mm), block size (BB), and dampening (λ\lambda). These are supposed to be tuned with respect to the model and available computational resources. In all of our runs, across all models and datasets, we use the same set of hyper-parameters which we found to work best for the BERTBASE \textrm{BERT}_{\textrm{BASE}}\,model on the SQuAD v1.1 dataset. We conjecture that further tuning for smaller models (3 and 6-layer models) could improve their results, but for simplicity and fairness to other methods, we apply the same ones found for the BERTBASE \textrm{BERT}_{\textrm{BASE}}\,.

Ablation studies. The procedure to find the optimal set of hyper-parameters for a model consists of a grid search over the possible hyper-parameter combinations and one-shot pruning runs to various high sparsity targets to evaluate the quality of the pruning approximation for each combination. We found that m=1024m=1024, B=50B=50, and λ=10−7\lambda=10^{-7} produce state-of-the-art results for a negligible computational overhead with the BERTBASE \textrm{BERT}_{\textrm{BASE}}\,model. Frantar et al. (2021) shows that larger block sizes require more gradients for better approximation. Given the massive size of the BERTBASE \textrm{BERT}_{\textrm{BASE}}\,model, we picked this setup as it was the best performing one that could still fit on a single 24GB RTX 3090 GPU card. In Figures 3, 4, and 5 we visualize a fraction of the one-shot pruning ablations with respect to all three hyper-parameters that motivated us to pick these specific values.

A.5 Downstream pruning

Teacher preparation. For all downstream pruning runs we make use of the KD from the fine-tuned BERTBASE \textrm{BERT}_{\textrm{BASE}}\,teacher outputs. The teacher is fine-tuned on the corresponding downstream task following the default hyper-parameters for SQuADhttps://github.com/huggingface/transformers/tree/main/ examples/pytorch/question-answering and GLUE (QQP and MNLI)https://github.com/huggingface/transformers/tree/main/ examples/pytorch/text-classification.

Pruning setup. In Table 9 we describe in detail all hyper-parameters for downstream pruning results presented in Tables 1 and 3. For easier comprehension, we also visualize learning rate schedules in Figures 6 and 8, and sparsity schedules in Figures 7 and 9.

3-, 6-layer models. We prepare our 3 and 6 layer models for downstream runs in two stages: layer dropping and retraining phase. We drop layers from our upstream teacher model (more details on it in Appendix A.6). After dropping, we retrain the remaining layers, following insights from Turc et al. (2019), in the same setup used to prepare the upstream teacher with addition of the KD from it.

A.6 Upstream pruning

Teacher preparation. We prepare a teacher for upstream pruning by following some insights from Liu et al. (2019). More concretely we start with the bert-base-uncasedhttps://huggingface.co/bert-base-uncased model, adopt pre-training on two datasets (BookCorpushttps://huggingface.co/datasets/bookcorpus & English Wikipediahttps://huggingface.co/datasets/wikipedia) with focus on the masked language modeling task (MLM) for 10-epochs with batch size 256 and learning rate linearly decaying to zero from the initial value of 1e-4.

Pruning setup. In Table 10 we describe in detail our upstream pruning recipe. As can be noticed, our upstream pruning recipe is just a downscaled version of our 30-epoch downstream-pruning recipe to 3-epochs.

A.7 Downstream quantization

We perform QAT on top of dense and 4-block pruned models on SQuAD v1.1 as shown in Table 3. We quantize to 8 bits the embedding matrices, linear modules of all encoder units which includes matrices in their attention and feed forward layers, and the linear module of the output layer. Weights that were pruned are kept constant (zero) during quantization (sparsity mask preserved). Non-linear operations within the Softmax, LayerNorm and GeLU are not quantized. For each dense and 4-block pruned model in Table 3, we perform a total of ten epochs training where the quantization observers are active for the first five and the remaining is fine-tuning. We do hyper-parameter search over the learning rates of 1e-4, 8e-5, 5e-5, 3e-5 and the distillation hardness of 0.9 and 1.0. We then pick the model with the best F1 score.

A.8 Additional performance metrics

Due to the space constraints, in the paper we report F1 score for SQuAD v1.1, matched accuracy for MNLI, and accuracy for QQP dataset. As all of our hyper-parameters for MNLI and QQP are exactly the same, we refer to these two datasets as GLUE. In Table 12 we report the additional metrics too: exact match (EM) for SQuAD v1.1, mismatched accuracy for MNLI, and F1 score for QQP dataset. Tables 15 and 16 present standard deviations of the corresponding results in Tables 1, 2 and 12. Finally, Table 14 presents the exact-match metric for the corresponding results in Table 3.

A.9 Inference speedups and compression ratios of compressed models

Details on the results shown in Figure 2 are drawn from Table 17. As shown in the results, not all compound compressed models yield improvements in inference or compression relative to retained model performance but those that do allow for massive improvements.

A.10 Responsible NLP Research - Reproducibility Checklist

In addition to many items from the “Reproducibility Checklist” which are already carefully addressed throughout the paper and Appendix sections, here we provide the remaining details to facilitate reproducibility of our results.

Datasets. Our experiments use existing and well established benchmarks for pre-training and fine-tuning of LLMs. Each dataset was used without any additional forms of modifications. Given that we did not modify any of the datasets, we did not inspect for personal, sensitive, or offensive content, nor did we perform any kind of anonymization. For pre-training, we make use of the Toronto Book Corpus (TBC) Zhu et al. (2015) https://huggingface.co/datasets/bookcorpus and the wikipedia.20200501.en Foundation https://huggingface.co/datasets/wikipedia. For fine-tuning we make use of SQuAD v1.1 Rajpurkar et al. (2016) https://huggingface.co/datasets/squad, Quora Duplicate Question Dataset (QQP) Shankar (2017) https://huggingface.co/datasets/glue, and Multi-Genre Natural Language Inference (MNLI) Williams et al. (2018) https://huggingface.co/datasets/glue datasets. All these datasets are publicly available via HuggingFace datasets repository Lhoest et al. (2021). The terms of usage and further details on each dataset can be found in their respective repositories.

Models. The model used as a starting point for all of our experiments is BERTBASE \textrm{BERT}_{\textrm{BASE}}\,, publicly available via HuggingFace Hub https://huggingface.co/bert-base-uncased. All other models presented in this paper will be released in openly-available repositories along with their compression recipes, training metrics and hyper-parameters.

A.10.2 Dataset Statistics

Dataset statistics are detailed in Table 18.

A.10.3 Computational Experiments

Upstream. All upstream runs are in general computationally expensive due to the large batch sizes and huge datasets. In our experiments we make use of 4x A100 40GB NVIDIA GPUs. In this configuration, a single training epoch takes approximately 6 hours. Since the cost of such a large compute instance is high, these experiments were only run with a single seed and without major hyper-parameter exploration.

Downstream. Our downstream experiments make use of various different GPU cards that were at out disposal: 16GB V100, 11GB RTX 2080 Ti, and 24GB RTX 3090. Each training epoch takes approximately 30 minutes, and as a result the 30 epoch runs take approximately 15 hours. For these experiments, we report mean results of three runs with different random seeds.

DeepSparse inference. We pair our compressed models with DeepSparse NeuralMagic (2021) a publicly-available sparsity-aware CPU inference engine. This CPU runtime can leverage both structured and unstructured sparsity, and quantization to deliver high performance on commodity CPUs. We ran DeepSparse on a 24-core Intel AWS c5.12xlarge server with 24 cores, 96 vCPUs, 192 GB of RAM and an AVX-512 compatible instruction set. All models are exported using the standard ONNXhttps://onnx.ai/ format.

A.10.4 Computational Packages

Our experiments build on publicly available libraries to ensure ease of reproduction and extensibility. All of our implementations, training and evaluation code are built on top of HuggingFace’s Transformers https://github.com/huggingface/transformers and Datasets https://github.com/huggingface/datasets libraries, NeuralMagic’s SparseML https://github.com/neuralmagic/sparseml library for model compression, and their DeepSparse https://github.com/neuralmagic/deepsparse engine for efficient inference on commodity CPUs.