Interpretable Neural Predictions with Differentiable Binary Variables
Jasmijn Bastings, Wilker Aziz, Ivan Titov
Introduction
Neural networks are bringing incredible performance gains on text classification tasks (Howard and Ruder, 2018; Peters et al., 2018; Devlin et al., 2019). However, this power comes hand in hand with a desire for more interpretability, even though its definition may differ (Lipton, 2016). While it is useful to obtain high classification accuracy, with more data available than ever before it also becomes increasingly important to justify predictions. Imagine having to classify a large collection of documents, while verifying that the classifications make sense. It would be extremely time-consuming to read each document to evaluate the results. Moreover, if we do not know why a prediction was made, we do not know if we can trust it.
What if the model could provide us the most important parts of the document, as a justification for its prediction? That is exactly the focus of this paper. We use a setting that was pioneered by Lei et al. (2016). A rationale is defined to be a short yet sufficient part of the input text; short so that it makes clear what is most important, and sufficient so that a correct prediction can be made from the rationale alone. One neural network learns to extract the rationale, while another neural network, with separate parameters, learns to make a prediction from just the rationale. Lei et al. model this by assigning a binary Bernoulli variable to each input word. The rationale then consists of all the words for which a 1 was sampled. Because gradients do not flow through discrete samples, the rationale extractor is optimized using REINFORCE (Williams, 1992). An regularizer is used to make sure the rationale is short.
We propose an alternative to purely discrete selectors for which gradient estimation is possible without REINFORCE, instead relying on a reparameterization of a random variable that exhibits both continuous and discrete behavior (Louizos et al., 2017). To promote compact rationales, we employ a relaxed form of regularization (Louizos et al., 2017), penalizing the objective as a function of the expected proportion of selected text. We also propose the use of Lagrangian relaxation to target a specific rate of selected input text.
Our contributions are summarized as follows:Code available at https://github.com/bastings/interpretable_predictions.
we present a differentiable approach to extractive rationales (§2) including an objective that allows for specifying how much text is to be extracted (§4);
we introduce HardKuma (§3), which gives support to binary outcomes and allows for reparameterized gradient estimates;
we empirically show that our approach is competitive with previous work and that HardKuma has further applications, e.g. in attention mechanisms. (§6).
Latent Rationale
We are interested in making NN-based text classifiers interpretable by (i) uncovering which parts of the input text contribute features for classification, and (ii) basing decisions on only a fraction of the input text (a rationale). Lei et al. (2016) approached (i) by inducing binary latent selectors that control which input positions are available to an NN encoder that learns features for classification/regression, and (ii) by regularising their architectures using sparsity-inducing penalties on latent assignments. In this section we put their approach under a probabilistic light, and this will then more naturally lead to our proposed method.
In text classification, an input is mapped to a distribution over target labels:
where we have a neural network architecture parameterize the model— collectively denotes the parameters of the NN layers in . That is, an NN maps from data space (e.g. sentences, short paragraphs, or premise-hypothesis pairs) to the categorical parameter space (i.e. a vector of class probabilities). For the sake of concreteness, consider the input a sequence . A target is typically a categorical outcome, such as a sentiment class or an entailment decision, but with an appropriate choice of likelihood it could also be a numerical score (continuous or integer).
Lei et al. (2016) augment this model with a collection of latent variables which we denote by . These variables are responsible for regulating which portions of the input contribute with predictors (i.e. features) to the classifier. The model formulation changes as follows:
where an NN predicts a sequence of Bernoulli parameters—one per latent variable—and the classifier is modified such that indicates whether or not is available for encoding. We can think of the sequence as a binary gating mechanism used to select a rationale, which with some abuse of notation we denote by . Figure 1 illustrates the approach.
Parameter estimation for this model can be done by maximizing a lower bound on the log-likelihood of the data derived by application of Jensen’s inequality:This can be seen as variational inference (Jordan et al., 1999) where we perform approximate inference using a data-dependent prior .
These latent rationales approach the first objective, namely, uncovering which parts of the input text contribute towards a decision. However note that an NN controls the Bernoulli parameters, thus nothing prevents this NN from selecting the whole of the input, thus defaulting to a standard text classifier. To promote compact rationales, Lei et al. (2016) impose sparsity-inducing penalties on latent selectors. They penalise for the total number of selected words, in (4), as well as, for the total number of transitions, fused lasso in (4), and approach the following optimization problem
via gradient-based optimisation, where and are fixed hyperparameters. The objective is however intractable to compute, the lowerbound, in particular, requires marginalization of binary sequences. For that reason, Lei et al. sample latent assignments and work with gradient estimates using REINFORCE (Williams, 1992).
The key ingredients are, therefore, binary latent variables and sparsity-inducing regularization, and therefore the solution is marked by non-differentiability. We propose to replace Bernoulli variables by rectified continuous random variables (Socci et al., 1998), for they exhibit both discrete and continuous behaviour. Moreover, they are amenable to reparameterization in terms of a fixed random source (Kingma and Welling, 2014), in which case gradient estimation is possible without REINFORCE. Following Louizos et al. (2017), we exploit one such distribution to relax regularization and thus promote compact rationales with a differentiable objective. In section 3, we introduce this distribution and present its properties. In section 4, we employ a Lagrangian relaxation to automatically target a pre-specified selection rate. And finally, in section 5 we present an example for sentiment classification.
Hard Kumaraswamy Distribution
Key to our model is a novel distribution that exhibits both continuous and discrete behaviour, in this section we introduce it. With non-negligible probability, samples from this distribution evaluate to exactly or exactly . In a nutshell: i) we start from a distribution over the open interval (see dashed curve in Figure 2); ii) we then stretch its support from to in order to include and (see solid curve in Figure 2); finally, iii) we collapse the probability mass over the interval to , and similarly, the probability mass over the interval to (shaded areas in Figure 2). This stretch-and-rectify technique was proposed by Louizos et al. (2017), who rectified samples from the BinaryConcrete (or GumbelSoftmax) distribution (Maddison et al., 2017; Jang et al., 2017). We adapted their technique to the Kumaraswamy distribution motivated by its close resemblance to a Beta distribution, for which we have stronger intuitions (for example, its two shape parameters transit rather naturally from unimodal to bimodal configurations of the distribution). In the following, we introduce this new distribution formally.We use uppercase letters for random variables (e.g. , , and ) and lowercase for assignments (e.g. , , ). For a random variable , is the probability density function (pdf), conditioned on parameters , and is the cumulative distribution function (cdf).
The Kumaraswamy is a close relative of the Beta distribution, though not itself an exponential family, with a simple cdf whose inverse
for , can be used to obtain samples
by transformation of a uniform random source . We can use this fact to reparameterize expectations (Nalisnick and Smyth, 2016).
2 Rectified Kumaraswamy
We stretch the support of the Kumaraswamy distribution to include and . The resulting variable takes on values in the open interval where and , with cdf
We now define a rectified random variable, denoted by , by passing a sample through a hard-sigmoid, i.e. . The resulting variable is defined over the closed interval $t=0h=0t\in(l,0]\operatorname{Kuma}(t|a,b,l,r)$ is available in closed form:
That is because all negative values of are deterministically mapped to zero. Similarly, samples are all deterministically mapped to , whose total mass amounts to
See Figure 2 for an illustration, and Appendix A for the complete derivations.
3 Reparameterization and gradients
Because this rectified variable is built upon a Kumaraswamy, it admits a reparameterisation in terms of a uniform variable . We need to first sample a uniform variable in the open interval and transform the result to a Kumaraswamy variable via the inverse cdf (10a), then shift and scale the result to cover the stretched support (10b), and finally, apply the rectifier in order to get a sample in the closed interval $$ (10c).
We denote this for short. Note that this transformation has two discontinuity points, namely, and . Though recall, the probability of sampling exactly or exactly is zero, which essentially means stochasticity circumvents points of non-differentiability of the rectifier (see Appendix A.3).
Controlled Sparsity
Following Louizos et al. (2017), we relax non-differentiable penalties by computing them on expectation under our latent model . In addition, we propose the use of Lagrangian relaxation to target specific values for the penalties. Thanks to the tractable Kumaraswamy cdf, the expected value of is known in closed form
In both cases, we make the assumption that latent variables are independent given , in Appendix B.1.2 we discuss how to estimate the regularizers for a model that conditions on the prefix of sampled assignments.
We can use regularizers to promote sparsity, but just how much text will our final model select? Ideally, we would target specific values and solve a constrained optimization problem. In practice, constrained optimisation is very challenging, thus we employ Lagrangian relaxation instead:
where is a vector of regularisers, e.g. expected and expected fused lasso, and is a vector of Lagrangian multipliers . Note how this differs from the treatment of Lei et al. (2016) shown in (4) where regularizers are computed for assignments, rather than on expectation, and where are fixed hyperparameters.
Sentiment Classification
As a concrete example, consider the case of sentiment classification where is a sentence and is a -way sentiment class (from very negative to very positive). The model consists of
where the shape parameters , i.e. two sequences of strictly positive scalars, are predicted by a NN, and the support boundaries are fixed hyperparameters.
We first specify an architecture that parameterizes latent selectors and then use a reparameterized sample to restrict which parts of the input contribute encodings for classification:We describe architectures using blocks denoted by , boldface letters for vectors, and the shorthand for a sequence .
where is an embedding layer, is a bidirectional encoder, and are feed-forward transformations with outputs, and turns the uniform sample into the latent selector (see §3). We then use the sampled to modulate inputs to the classifier:
where and are recurrent cells such as LSTMs (Hochreiter and Schmidhuber, 1997) that process the sequence in different directions, and is a feed-forward transformation with output. Note how modulates features of the input that are available to the recurrent composition function.
We then obtain gradient estimates of via Monte Carlo (MC) sampling from
where is a shorthand for elementwise application of the transformation from uniform samples to HardKuma samples. This reparameterisation is the key to gradient estimation through stochastic computation graphs (Kingma and Welling, 2014; Rezende et al., 2014).
At test time we make predictions based on what is the most likely assignment for each . We across configurations of the distribution, namely, , , or . When the continuous interval is more likely, we take the expected value of the underlying Kumaraswamy variable.
Experiments
We perform experiments on multi-aspect sentiment analysis to compare with previous work, as well as experiments on sentiment classification and natural language inference. All models were implemented in PyTorch, and Appendix B provides implementation details.
When rationalizing predictions, our goal is to perform as well as systems using the full input text, while using only a subset of the input text, leaving unnecessary words out for interpretability.
1 Multi-aspect Sentiment Analysis
In our first experiment we compare directly with previous work on rationalizing predictions (Lei et al., 2016). We replicate their setting.
A pre-processed subset of the BeerAdvocatehttps://www.beeradvocate.com/ data set is used (McAuley et al., 2012). It consists of 220,000 beer reviews, where multiple aspects (e.g. look, smell, palate) are rated. As shown in Figure 1, a review typically consists of multiple sentences, and contains a 0-5 star rating (e.g. 3.5 stars) for each aspect. Lei et al. mapped the ratings to scalars in $$.
We use the models described in §5 with two small modifications: 1) since this is a regression task, we use a sigmoid activation in the output layer of the classifier rather than a softmax,From a likelihood learning point of view, we would have assumed a Logit-Normal likelihood, however, to stay closer to Lei et al. (2016), we employ mean squared error. and 2) we use an extra RNN to condition on :
For a fair comparison we follow Lei et al. by using RCNNAn RCNN cell can replace any LSTM cell and works well on text classification problems. See appendix B. cells rather than LSTM cells for encoding sentences on this task. Since this cell is not widely used, we verified its performance in Table 1. We observe that the BiRCNN performs on par with the BiLSTM (while using 50% fewer parameters), and similarly to previous results.
A test set with sentence-level rationale annotations is available. The precision of a rationale is defined as the percentage of words with that is part of the annotation. We also evaluate the predictions made from the rationale using mean squared error (MSE).
For our baseline we reimplemented the approach of Lei et al. (2016) which we call Bernoulli after the distribution they use to sample from. We also report their attention baseline, in which an attention score is computed for each word, after which it is simply thresholded to select the top-k percent as the rationale.
Table 2 shows the precision and the percentage of selected words for the first three aspects. The models here have been selected based on validation MSE and were tuned to select a similar percentage of words (‘selected’). We observe that our Bernoulli reimplementation reaches the precision similar to previous work, doing a little bit worse for the ‘look’ aspect. Our HardKuma managed to get even higher precision, and it extracted exactly the percentage of text that we specified (see §4).We tried to use Lagrangian relaxation for the Bernoulli model, but this led to instabilities (e.g. all words selected). Figure 3 shows the MSE for all aspects for various percentages of extracted text. We observe that HardKuma does better with a smaller percentage of text selected. The performance becomes more similar as more text is selected.
2 Sentiment Classification
We also experiment on the Stanford Sentiment Treebank (SST) (Socher et al., 2013). There are 5 sentiment classes: very negative, negative, neutral, positive, and very positive. Here we use the HardKuma model described in §5, a Bernoulli model trained with REINFORCE, as well as a BiLSTM.
Figure 4 shows the classification accuracy for various percentages of selected text. We observe that HardKuma outperforms the Bernoulli model at each percentage of selected text. HardKuma reaches full-text baseline performance already around 40% extracted text. At that point, it obtains a test score of 45.84, versus 42.22 for Bernoulli and 47.4 for the full-text baseline.
We wonder what kind of words are dropped when we select smaller amounts of text. For this analysis we exploit the word-level sentiment annotations in SST, which allows us to track the sentiment of words in the rationale. Figure 5 shows that a large portion of dropped words have neutral sentiment, and it seems plausible that exactly those words are not important features for classification. We also see that HardKuma drops (relatively) more neutral words than Bernoulli.
3 Natural Language Inference
In Natural language inference (NLI), given a premise sentence and a hypothesis sentence , the goal is to predict their relation which can be contradiction, entailment, or neutral. As our dataset we use the Stanford Natural Language Inference (SNLI) corpus (Bowman et al., 2015).
We use the Decomposable Attention model (DA) of Parikh et al. (2016).Better results e.g. Chen et al. (2017) and data sets for NLI exist, but are not the focus of this paper. DA does not make use of LSTMs, but rather uses attention to find connections between the premise and the hypothesis that are predictive of the relation. Each word in the premise attends to each word in the hypothesis, and vice versa, resulting in a set of comparison vectors which are then aggregated for a final prediction. If there is no link between a word pair, it is not considered for prediction.
Because the premise and hypothesis interact, it does not make sense to extract a rationale for the premise and hypothesis independently. Instead, we replace the attention between premise and hypothesis with HardKuma attention. Whereas in the baseline a similarity matrix is -normalized across rows (premise to hypothesis) and columns (hypothesis to premise) to produce attention matrices, in our model each cell in the attention matrix is sampled from a HardKuma parameterized by . To promote sparsity, we use the relaxed to specify the desired percentage of non-zero attention cells. The resulting matrix does not need further normalization.
With a target rate of 10%, the HardKuma model achieved 8.5% non-zero attention. Table 3 shows that, even with so many zeros in the attention matrices, it only does about 1% worse compared to the DA baseline. Figure 6 shows an example of HardKuma attention, with additional examples in Appendix B. We leave further explorations with HardKuma attention for future work.
Related Work
This work has connections with work on interpretability, learning from rationales, sparse structures, and rectified distributions. We discuss each of those areas.
Machine learning research has been focusing more and more on interpretability (Gilpin et al., 2018). However, there are many nuances to interpretability (Lipton, 2016), and amongst them we focus on model transparency.
One strategy is to extract a simpler, interpretable model from a neural network, though this comes at the cost of performance. For example, Thrun (1995) extract if-then rules, while Craven and Shavlik (1996) extract decision trees.
There is also work on making word vectors more interpretable. Faruqui et al. (2015) make word vectors more sparse, and Herbelot and Vecchi (2015) learn to map distributional word vectors to model-theoretic semantic vectors.
Similarly to Lei et al. (2016), Titov and McDonald (2008) extract informative fragments of text by jointly training a classifier and a model predicting a stochastic mask, while relying on Gibbs sampling to do so. Their focus is on using the sentiment labels as a weak supervision signal for opinion summarization rather than on rationalizing classifier predictions.
There are also related approaches that aim to interpret an already-trained model, in contrast to Lei et al. (2016) and our approach where the rationale is jointly modeled. Ribeiro et al. (2016) make any classifier interpretable by approximating it locally with a linear proxy model in an approach called LIME, and Alvarez-Melis and Jaakkola (2017) propose a framework that returns input-output pairs that are causally related.
Our work is different from approaches that aim to improve classification using rationales as an additional input (Zaidan et al., 2007; Zaidan and Eisner, 2008; Zhang et al., 2016). Instead, our rationales are latent and we are interested in uncovering them. We only use annotated rationales for evaluation.
Also arguing for enhanced interpretability, Niculae and Blondel (2017) propose a framework for learning sparsely activated attention layers based on smoothing the operator. They derive a number of relaxations to , including itself, but in particular, they target relaxations such as (Martins and Astudillo, 2016) which, unlike , are sparse (i.e. produce vectors of probability values with components that evaluate to exactly ). Their activation functions are themselves solutions to convex optimization problems, to which they provide efficient forward and backward passes. The technique can be seen as a deterministic sparsely activated layer which they use as a drop-in replacement to standard attention mechanisms. In contrast, in this paper we focus on binary outcomes rather than -valued ones. Niculae et al. (2018) extend the framework to structured discrete spaces where they learn sparse parameterizations of discrete latent models. In this context, parameter estimation requires exact marginalization of discrete variables or gradient estimation via REINFORCE. They show that oftentimes distributions are sparse enough to enable exact marginal inference.
Peng et al. (2018) propose SPIGOT, a proxy gradient to the non-differentiable operator. This proxy requires an solver (e.g. Viterbi for structured prediction) and, like the straight-through estimator (Bengio et al., 2013), is a biased estimator. Though, unlike ST it is efficient for structured variables. In contrast, in this work we chose to focus on unbiased estimators.
The idea of rectified distributions has been around for some time. The rectified Gaussian distribution (Socci et al., 1998), in particular, has found applications to factor analysis (Harva and Kaban, 2005) and approximate inference in graphical models (Winn and Bishop, 2005). Louizos et al. (2017) propose to stretch and rectify samples from the BinaryConcrete (or GumbelSoftmax) distribution (Maddison et al., 2017; Jang et al., 2017). They use rectified variables to induce sparsity in parameter space via a relaxation to . We adapt their technique to promote sparse activations instead. Rolfe (2017) learns a relaxation of a discrete random variable based on a tractable mixture of a point mass at zero and a continuous reparameterizable density, thus enabling reparameterized sampling from the half-closed interval . In contrast, with HardKuma we focused on giving support to both 0s and 1s.
Conclusions
We presented a differentiable approach to extractive rationales, including an objective that allows for specifying how much text is to be extracted. To allow for reparameterized gradient estimates and support for binary outcomes we introduced the HardKuma distribution. Apart from extracting rationales, we showed that HardKuma has further potential uses, which we demonstrated on premise-hypothesis attention in SNLI. We leave further explorations for future work.
Acknowledgments
We thank Luca Falorsi for pointing us to Louizos et al. (2017), which inspired the HardKumaraswamy distribution. This work has received funding from the European Research Council (ERC StG BroadSem 678254), the European Union’s Horizon 2020 research and innovation programme (grant agreement No 825299, GoURMET), and the Dutch National Science Foundation (NWO VIDI 639.022.518, NWO VICI 277-89-002).
References
Appendix A Kumaraswamy distribution
A Kumaraswamy-distributed variable takes on values in the open interval and has density
We can generalise the support of a Kumaraswamy variable by specifying two constants and transforming a random variable to obtain as shown in (20, left).
where by definition. This affine transformation leaves the cdf unchanged, i.e.
Thus we can obtain samples from this generalised-support Kumaraswamy by sampling from a uniform distribution , applying the inverse transform (19), then shifting and scaling the sample according to (20, left).
A.2 Rectified Kumaraswamy
First, we stretch a Kumaraswamy distribution to include and in its support, that is, with and , we define . Then we apply a hard-sigmoid transformation to this variable, that is, , which results in a rectified distribution which gives support to the closed interval $$. We denote this rectified variable by
is the probability of sampling exactly , where
is the probability of sampling exactly 1, and
is the probability of drawing a continuous value in . Note that we used the result in (22) to express these probabilities in terms of the tractable cdf of the original Kumaraswamy variable.
A.3 Reparameterized gradients
Let us consider the case where we need derivatives of a function of the underlying uniform variable , as when we compute reparameterized gradients in variational inference. Note that
by chain rule. The term depends on a differentiable observation model and poses no challenge; the term is the derivative of the hard-sigmoid function, which is for or , for , and undefined for ; the term follows directly from (20, left); the term depends on the Kumaraswamy inverse cdf (19) and also poses no challenge. Thus the only two discontinuities happen for , which is a measure set under the stretched Kumaraswamy: we say this reparameterisation is differentiable almost everywhere, a useful property which essentially circumvents the discontinuity points of the rectifier.
A.4 HardKumaraswamy PDF and CDF
Figure 8 plots the pdf of the HardKumaraswamy for various a and b parameters. Figure 9 does the same but with the cdf.
Appendix B Implementation Details
Our hyperparameters are taken from Lei et al. (2016) and listed in Table 4. The pre-trained word embeddings and data sets are available online at http://people.csail.mit.edu/taolei/beer/. We train for 100 epochs and select the best models based on validation loss. For the MSE trade-off experiments on all aspects combined, we train for a maximum of 50 epochs.
For the Bernoulli baselines we vary weight among , just as in the original paper. We set the fused lasso (coherence) weight to .
For the HardKuma models we set a target selection rate to the values targeted in Table 2, and optimize to this end using the Lagrange multiplier. We chose the fused lasso weight from .
In our multi-aspect sentiment analysis experiments we use the RCNN of Lei et al. (2016). Intuitively, the RCNN is supposed to capture n-gram features that are not necessarily consecutive. We use the bigram version (filter width ) used in Lei et al. (2016), which is defined as:
B.1.2 Expected values for dependent latent variables
The expected is a chain of nested expectations, and we solve each term
as a function of a sampled prefix, and the shape parameters are predicted in sequence.
B.2 Sentiment Classification (SST)
For sentiment classification we make use of the PyTorch bidirectional LSTM module for encoding sentences, for both the rationale extractor and the classifier. The BiLSTM final states are concatenated, after which a linear layer followed by a softmax produces the prediction. Hyperparameters are listed in Table 5. We apply dropout to the embeddings and to the input of the output layer.
B.3 Natural Language Inference (SNLI)
Our hyperparameters are taken from Parikh et al. (2016) and listed in Table 6. Different from Parikh et al. is that we use Adam as the optimizer and a batch size of 64. Word embeddings are projected to 200 dimensions with a trained linear layer. Unknown words are mapped to 100 unknown word classes based on the MD5 hash function, just as in Parikh et al. (2016), and unknown word vectors are randomly initialized. We train for 100 epochs, evaluate every 1000 updates, and select the best model based on validation loss. Figure 10 shows a correct and incorrect example with HardKuma attention for each relation type (entailment, contradiction, neutral).