Low-Complexity Probing via Finding Subnetworks
Steven Cao, Victor Sanh, Alexander M. Rush
Introduction
While pre-training has produced large gains for natural language tasks, it is unclear what a model learns during pre-training. Research in probing investigates this question by training a shallow classifier on top of the pre-trained model’s internal representations to predict some linguistic property (Adi et al. 2016; Shi et al. 2016; Tenney et al. 2019, inter alia). The resulting accuracy is then roughly indicative of the model encoding that property.
However, it is unclear how much is learned by the probe versus already captured in the model representations. This question has been the subject of much recent debate (Hewitt and Liang 2019; Voita and Titov 2020; Pimentel et al. 2020b, inter alia). We would like the probe to find only and all properties captured by a model, leading to a tradeoff between accuracy and complexity: a linear probe is insufficient to find the non-linear patterns in neural models, but a deeper multi-layer perceptron (MLP) is complex enough to learn the task on its own.
Motivated by this tradeoff and the goal of low-complexity probes, we consider a different approach based on pruning. Specifically, we search for a subnetwork — a version of the model with a subset of the weights set to zero — that performs the task of interest. As our search procedure, we build upon past work in pruning and perform gradient descent on a continuous relaxation of the search problem (Louizos et al. 2017; Mallya et al. 2018; Sanh et al. 2020). The resulting probe has many fewer free parameters than MLP probes.
Our experiments evaluate the accuracy-complexity tradeoff compared to MLP probes on an array of linguistic tasks. First, we find that the neuron subnetwork probe has both higher accuracy on pre-trained models and lower accuracy on random models, so it is both better at finding properties of interest and less able to learn the tasks on its own. Next, we measure complexity as the bits needed to transmit the probe parameters (Pimentel et al. 2020a; Voita and Titov 2020). Varying the complexity of each probe, we find that subnetwork probing Pareto-dominates MLP probing in that it achieves higher accuracy given any desired complexity. Finally, we analyze the resulting subnetworks across various tasks and find that lower-level tasks are captured in lower layers, reproducing similar findings in past work (Tenney et al. 2019). These results suggest that subnetwork probing is an effective new direction for improving our understanding of pre-training.
Related Work
Probing investigates whether a model captures some hypothesized property and typically involves learning a shallow classifier on top of the model’s frozen internal representations (Adi et al. 2016; Shi et al. 2016; Conneau et al. 2018). Recent work has primarily applied this technique to pre-trained models. While probing is also used in other domains (e.g. neural decoding), we focus on understanding neural models. Therefore, one source of strength for our probe is that we exploit the entire model, rather than only operating on representations. Clark et al. 2019, Hewitt and Manning 2019, and Manning et al. 2020 found that BERT captures various properties of syntax. Tenney et al. 2019 probed the layers of BERT for an array of tasks, and they found that their localization mirrored the classical NLP pipeline (part-of-speech, parsing, named entity recognition, semantic roles, coreference) in that lower-level tasks were captured in the lower layers.
However, these results are difficult to interpret due to the use of a learned classifier. One line of work suggests comparing the probe accuracy to random baselines, e.g. random models (Zhang and Bowman 2018) or random control tasks (Hewitt and Liang 2019). Other works take an information-theoretic view: Voita and Titov 2020 measure the complexity of the probe in terms of the bits needed to transmit its parameters, while Pimentel et al. 2020b argue that probing should measure mutual information between the representation and the property. Pimentel et al. 2020a propose a Pareto approach where they plot accuracy versus probe complexity, unifying several of these goals. We use these proposed metrics to compare our probing method to standard probing approaches.
Subnetworks.
While pruning is widely used for model compression, some works have explored pruning as a technique for learning as well. Mallya et al. 2018 found that a model trained on ImageNet could be used for new tasks by learning a binary mask over the weights. More recently, Radiya-Dixit and Wang 2020 and Zhao et al. 2020 showed the analogous result in NLP that weight pruning can be used as an alternative to fine-tuning for pre-trained models. Our paper seeks to use pruning to reveal what the model already captures, rather than learn new tasks.
Subnetwork Probing
Given a task and a pre-trained encoder model with a classification head, our goal is to find a subnetwork with high accuracy on that task, where a subnetwork is the model with a subset of the encoder weights masked, i.e. set to zero. We search for this subnetwork via supervised gradient descent on the head and a continuous relaxation of the mask. We also mask at several levels of granularity, including pruning weights, neurons, or layers.
where denotes the sigmoid and , are constants. This random variable can be thought of as a soft version of the Bernoulli. follows the concrete (or Gumbel-Softmax) distribution with temperature (Maddison et al. 2016; Jang et al. 2016). To put non-zero mass on and , the distribution is stretched to the interval and clamped back to $$.
We will denote the mask as and the masked weights as . We can then optimize the mask parameters via gradient descent. Specifically, let denote the model. Then, given a data point and a loss function , we can minimize the expectation of the loss, or
We estimate the expectation via sampling: we sample a single and take the gradient . To encourage sparsity, we penalize the mask based on the probability it is non-zero, or
Letting denote regularization strength, our objective becomes . Departing from past work, we schedule linearly to improve search: it stays fixed at for the first 25% of training, linearly increases to for the next 50%, and then stays fixed. We set in our evaluation experiments.
Probe Evaluation
To evaluate the accuracy-complexity tradeoff of a probe, we adapt methodology from recent work. First, we consider the non-parametric test of probing a random model (Zhang and Bowman 2018). We check probe accuracy on the pre-trained model, the model with the encoder randomly reset (reset encoder), and the model with the encoder and embeddings reset (reset all). An ideal probe should achieve high accuracy on the pre-trained model and low accuracy on the reset models. The reset encoder model contains some non-contextual information from its word embeddings, but no modeling of context; therefore, we would expect it to have better probe accuracy on tasks based mainly on word type (e.g. part-of-speech tagging).
Next, we consider a parametric test based on probe complexity. We first vary the complexity of each probe, where for subnetwork probing we associate multiple encoder weights with a single mask, For subnetworks, the pre-trained model has 72 matrices of size ; see https://github.com/huggingface/transformers/blob/v3.4.0/src/transformers/modeling_bert.py. For each matrix, let and denote the number of rows and columns per mask. Then, we set to , , , , , , , , and . corresponds to masking entire matrices, to masking neurons, and to masking weights. and for the MLP probe we restrict the rank of the hidden layer. We then plot the resulting accuracy-complexity curve (Pimentel et al. 2020a).
To plot this curve, we need a measure of complexity that can compare probes of different types. Therefore, we measure complexity as the number of bits needed to transmit the probe parameters (Voita and Titov 2020), where for simplicity we use a uniform encoding. In the subnetwork case, this encoding corresponds to using a single bit for each mask parameter. In the case of an MLP probe, each parameter is a real number, so the number of bits per parameter depends on its range and precision. For example, if each parameter lies in and requires precision, then we need bits per parameter. To avoid having the choice of precision impact results, we plot lower and upper bounds of and bits per parameter.
Experimental Setup
We probe bert-base-uncased (Devlin et al. 2019; Wolf et al. 2020) for the following tasks:
(1) Part-of-speech Tagging: We use the part-of-speech tags in the universal dependencies dataset (Zeman et al. 2017). As our classification head, we use dropout with probability , followed by a linear layer and softmax projecting from the BERT dimension to the number of tags.
(2) Dependency Parsing: We use the universal dependencies dataset (Zeman et al. 2017) and the biaffine head for classification (Dozat and Manning 2016). We report macro-averaged labeled attachment score.
(3) Named Entity Recognition (NER): We use the data from the CoNLL 2003 shared task (Tjong Kim Sang and De Meulder 2003) and the same classification head as for part-of-speech tagging. We report F1 using the CoNLL 2003 script.
Our primary probing baseline is the MLP probe with one hidden layer (MLP-1):
While we vary the complexity of each probe to produce the accuracy-complexity plot, we default to neuron subnetwork probing and full rank MLP-1 probing in all other experiments.
Results
Table 1 shows the results from the non-parametric experiments. When probing the pre-trained model, the subnetwork probe has much higher accuracy than the MLP-1 probe across all tasks. Furthermore, when probing the random models, the subnetwork probe has much lower accuracy for dependency parsing and NER, suggesting that the probe is less able to learn the task on its own. Overall, these numbers suggest that the subnetwork probe is a more faithful probe in that it finds properties when they are present, and does not find them in a random model.
Figure 1 plots the results from the parametric experiments, where we vary the complexity of each probe, apply it to the pre-trained model, and plot the resulting accuracy-complexity curve. We find that the subnetwork probe Pareto-dominates the MLP-1 probe in that it achieves higher accuracy for any complexity, even if we assume an overly optimistic MLP-1 lower bound of 1 bit per parameter. In particular, for part-of-speech and dependency parsing, the subnetwork probe achieves high accuracy even when given only 72 bits, while the MLP-1 probe falls off heavily at 20K bits.
Subnetwork Analysis.
An auxiliary benefit of subnetwork probing is that we can examine the subnetworks produced by the procedure. One possibility is to look at the locations of the subnetworks, and one way to examine location is to count the number of unmasked weights in each layer. Figure 2 shows locations of the remaining parameters in the subnetworks extracted from the pre-trained model and the random encoder model. To prune as many parameters as possible, we set to be the largest out of such that accuracy is within 10% of fine-tuning accuracy (see the Appendix for more details). We then examine the sparsity levels of the attention heads for each layer. While reset encoder model’s subnetworks are uniformly distributed across the layers, the pre-trained model’s subnetworks are localized and follow the order part-of-speech dependencies NER, reproducing the order found in Tenney et al. 2019. While Tenney et al. 2019 derived layer importance by training classifiers at each layer, we find location directly via pruning. This experiment strengthens their result and represents one example where subnetwork probing reveals additional insights into the model beyond accuracy.
Conclusion
Together, these results show that subnetwork probing is more faithful to the model and offers richer analysis than existing probing approaches. While this work explores accuracy and location-based analysis, there are other possible directions, e.g., applying neuron explainability techniques. Therefore, we see subnetwork probing as a fruitful new direction for understanding pre-training.
Ethical Considerations
While pre-trained models have improved performance for many NLP tasks, they exhibit biases present in the pre-training corpora (Manzini et al. 2019; Tan and Celis 2019; Kurita et al. 2019, inter alia). As a result, deploying pre-trained models runs the risk of reinforcing social biases. Probing gives us a tool to better understand and hopefully mitigate these biases. As one example of such a study, Vig et al. 2020 analyze how neurons and attention heads contribute to gender bias in pre-trained transformers. Therefore, while we analyze linguistic tasks in our paper, our method could also provide insights into model bias, e.g. by analyzing subnetworks for bias detection tasks like CrowS-Pairs (Nangia et al. 2020) or StereoSet (Nadeem et al. 2020).
Acknowledgements
We would like to thank Eric Wallace, Kevin Yang, Ruiqi Zhong, Dan Klein, and Yacine Jernite for their useful comments and feedback. This work was done during an internship at Hugging Face.
References
Appendix A Appendix
The mask parameters are optimized using Adam with , , and learning rate with linear warmup for the first of the data. The classification head parameters are also optimized using Adam with the same hyperparameters and warmup, except with learning rate . The MLP-1 and fine-tuning baselines are also optimized using Adam with the same hyperparameters, warmup, and learning rate . We train for 30 epochs for all tasks.
A.2 Varying Regularization Strength
Table 2 shows probing accuracies for . Our method is consistently more selective than MLP-1 across the various values of , except for , which seems to require too much sparsity.
A.3 Reproducibility Checklist
Experiments were run in Google Colab using a single 12GB NVIDIA Tesla K80 GPU. For each task, one run of fine-tuning took about half an hour. We used the transformers implementation of the bert-base-uncased model (Wolf et al. 2020; Devlin et al. 2019), which has 12 layers, 768 hidden dimension, 12 heads, and 110M parameters. As data, we used the dev (2002 examples) and train (12541 examples) splits of the English universal dependencies dataset (Zeman et al. 2017), and the test (3235 examples) and train (13862 examples) splits of the CoNLL 2003 NER shared task (Tjong Kim Sang and De Meulder 2003).