Causal Proxy Models for Concept-Based Model Explanations

Zhengxuan Wu, Karel D'Oosterlinck, Atticus Geiger, Amir Zur, Christopher Potts

Introduction

The gold standard for explanation methods in AI should be to elucidate the causal role that a model’s representations play in its overall behavior – to truly explain why the model makes the predictions it does. Causal explanation methods seek to do this by resolving the counterfactual question of what the model would do if input XX were changed to a relevant counterfactual version X′X^{\prime}. Unfortunately, even though neural networks are fully observed, deterministic systems, we still encounter the fundamental problem of causal inference (Holland, 1986): for a given ground-truth input XX, we never observe the counterfactual inputs X′X^{\prime} necessary for isolating the causal effects of model representations on outputs. The issue is especially pressing in domains where it is hard to synthesize approximate counterfactuals. In response to this, explanation methods typically do not explicitly train on counterfactuals at all.

In this paper, we show that robust explanation methods for NLP models can be obtained using texts approximating true counterfactuals. The heart of our proposal is the Causal Proxy Model (CPM). CPMs are trained to mimic both the factual and counterfactual behavior of a black-box model N\mathcal{N}{}. We explore two different methods for training such explainers. These methods share a distillation-style objective that pushes them to mimic the factual behavior of N\mathcal{N}{}, but they differ in their counterfactual objectives. The input-based method CPMIN\text{CPM{}}_{\text{IN}} appends to the factual input a new token associated with the counterfactual concept value. The hidden-state method CPMHI\text{CPM{}}_{\text{HI}} employs the Interchange Intervention Training (IIT) method of Geiger et al. (2022) to localize information about the target concept in specific hidden states. Figure 1 provides a high-level overview.

We evaluate these methods on the CEBaB benchmark for causal explanation methods (Abraham et al., 2022), which provides large numbers of original examples (restaurant reviews) with human-created counterfactuals for specific concepts (e.g., service quality), with all the texts labeled for their concept-level and text-level sentiment. We consider two types of approximate counterfactuals derived from CEBaB: texts written by humans to approximate a specific counterfactual, and texts sampled using metadata-guided heuristics. Both approximate counterfactual strategies lead to state-of-the-art performance on CEBaB for both CPMIN\text{CPM{}}_{\text{IN}} and CPMHI\text{CPM{}}_{\text{HI}}.

We additionally identify two other benefits of using CPMs to explain models. First, both CPMIN\text{CPM{}}_{\text{IN}} and CPMHI\text{CPM{}}_{\text{HI}} have factual performance comparable to that of the original black-box model N\mathcal{N}{} and can explain their own behavior extremely well. Thus, the CPM for N\mathcal{N}{} can actually replace N\mathcal{N}{}, leading to more explainable deployed models. Second, CPMHI\text{CPM{}}_{\text{HI}} models localize concept-level information in their hidden representations, which makes their behavior on specific inputs very easy to explain. We illustrate this using Path Integrated Gradients (Sundararajan et al., 2017), which we adapt to allow input-level attributions to be mediated by the intermediate states that were targeted for localization. Thus, while both CPMIN\text{CPM{}}_{\text{IN}} and CPMHI\text{CPM{}}_{\text{HI}} are comparable as explanation methods according to CEBaB, the qualitative insights afforded by CPMHI\text{CPM{}}_{\text{HI}} models may given them the edge when it comes to explanations.

Related Work

Understanding model behavior serves many goals for large-scale AI systems, including transparency (Kim, 2015; Lipton, 2018; Pearl, 2019; Ehsan et al., 2021), trustworthiness (Ribeiro et al., 2016; Guidotti et al., 2018; Jacovi & Goldberg, 2020; Jakesch et al., 2019), safety (Amodei et al., 2016; Otte, 2013), and fairness (Hardt et al., 2016; Kleinberg et al., 2017; Goodman & Flaxman, 2017; Mehrabi et al., 2021). With CPMs, our goal is to achieve explanations that are causally motivated and concept-based, and so we concentrate here on relating existing methods to these two goals.

Feature attribution methods estimate the importance of features, generally by inspecting learned weights directly or by perturbing features and studying the effects this has on model behavior (Molnar, 2020; Ribeiro et al., 2016). Gradient-based feature attribution methods extend this general mode of explanation to the hidden representations in deep networks (Zeiler & Fergus, 2014; Springenberg et al., 2014; Binder et al., 2016; Shrikumar et al., 2017; Sundararajan et al., 2017). Concept Activation Vectors (CAVs; Kim et al. 2018; Yeh et al. 2020) can also be considered feature attribution methods, as they probe for semantically meaningful directions in the model’s internal representations and use these to estimate the importance of concepts on the model predictions. While some methods in this space do have causal interpretations (e.g., Sundararajan et al. 2017; Yeh et al. 2020), most do not. In addition, most of these methods offer explanations in terms of specific (sets of) features/neurons. (Methods based on CAVs operate directly in terms of more abstract concepts.)

Intervention-based methods study model representations by modifying them in systematic ways and observing the resulting model behavior. These methods are generally causally motivated and allow for concept-based explanations. Examples of methods in this space include causal mediation analysis (Vig et al., 2020; De Cao et al., 2021; Ban et al., 2022), causal effect estimation (Feder et al., 2020; Elazar et al., 2021; Abraham et al., 2022; Lovering & Pavlick, 2022), tensor product decomposition (Soulos et al., 2020), and causal abstraction analysis (Geiger et al., 2020; 2021). CPMs are most closely related to the method of IIT (Geiger et al., 2021), which extends causal abstraction analysis to optimization.

Probing is another important class of explanation method. Traditional probes do not intervene on the target model, but rather only seek to find information in it via supervised models (Conneau et al., 2018; Tenney et al., 2019) or unsupervised models (Clark et al., 2019; Manning et al., 2020; Saphra & Lopez, 2019). Probes can identify concept-based information, but they cannot offer guarantees that probed information is relevant for model behavior (Geiger et al., 2021). For causal guarantees, it is likely that some kind of intervention is required. For example, Elazar et al. (2021) and Feder et al. (2020) remove information from model representations to estimate the causal role of that information. Our CPMs employ a similar set of guiding ideas but are not limited to removing information.

Counterfactual explanation methods aim to explain model behavior by providing a counterfactual example that changes the model behavior (Goyal et al., 2019; Verma et al., 2020; Wu et al., 2021). Counterfactual explanation methods are inherently causal. If they can provide counterfactual examples with regard to specific concepts, they are also concept-based.

Some explanation methods train a model making explicit use of intermediate variables representing concepts. Manipulating these intermediate variables at inference time yields causal concept-based model explanations (Koh et al., 2020; Künzel et al., 2019).

Evaluating methods in this space has been a persistent challenge. In prior literature, explanation methods have often been evaluated against synthetic datasets (Feder et al., 2020; Yeh et al., 2020). In response, Abraham et al. (2022) introduced the CEBaB dataset, which provides a human-validated concept-based dataset to truthfully evaluate different causal concept-based model explanation methods. Our primary evaluations are conducted on CEBaB.

Causal Proxy Model (CPM)

Causal Proxy Models (CPMs) are causal concept-based explanation methods. Given a factual input xu,vx_{u,v} and a description of a concept intervention Ci←c′{C_{i}}\leftarrow{c^{\prime}}, they estimate the effect of the intervention on model output. The present section introduces our two core CPM variants in detail. We concentrate here on introducing the structure of these models and their objectives, and we save discussion of associated metrics for explanation methods for Section 4.

Our discussion is grounded in the causal model depicted in Figure 1(a), which aligns well with the CEBaB benchmark. Two exogenous variables UU and VV together represent the complete state of the world and generate some textual data XX. The effect of exogenous variable UU on the data XX is completely mediated by a set of intermediate variables C1,C2…,CkC_{1},C_{2}\dots,C_{k}, which we refer to as concepts. Therefore, we can think of UU as the part of the world that gives rise to these concepts {C}1k\{C\}_{1}^{k}.

Using this causal model, we can describe counterfactual data – data that arose under a counterfactual state of the world (right diagram in Figure 1(a)). Our factual text is xu,vx_{u,v}, and we use xu,vCi←c′x_{u,v}^{{C_{i}}\leftarrow{c^{\prime}}} for the counterfactual text obtained by intervening on concept CiC_{i} to set its value to c′c^{\prime}. The counterfactual xu,vCi←c′x_{u,v}^{{C_{i}}\leftarrow{c^{\prime}}} describes the output when the value of CiC_{i} is set to c′c^{\prime}, all else being held equal.

Approximate Counterfactuals

where xu,v;tCi←c′x_{u,v};t_{{C_{i}}\leftarrow{c^{\prime}}} in Eqn. 2 denotes the concatenation of the factual input and the token describing the intervention. CES\text{CE}_{\text{S}} represents the smoothed cross-entropy loss (Hinton et al., 2015), measuring the divergence between the output logits of both models. The objective in Eqn. 1 pushes P\mathcal{P} to predict the same output as N\mathcal{N}{} under conventional circumstances (Figure 1(c)), while Eqn. 2 pushes P\mathcal{P} to predict the counterfactual behavior of N\mathcal{N}{} when a descriptor of the intervention is given (Figure 1(d)).These objectives are described with regard to a single approximate counterfactual pair for the sake of clarity. At train-time, we aggregate the objective over all considered training pairs. We take CiC_{i}{} to always represent the intervened-upon concept. The weights of N\mathcal{N}{} are frozen.

At inference time, approximate counterfactuals are inaccessible. To explain model N\mathcal{N}{}, we append the newly learned descriptor tokens tCi←c′t_{{C_{i}}\leftarrow{c^{\prime}}} to a factual input, upon which P\mathcal{P} predicts a counterfactual output for this input, used to estimate the counterfactual behavior of N\mathcal{N}{} under this intervention.

Our CPMHI\text{CPM{}}_{\text{HI}} models are trained on the same data and with the same set of goals as CPMIN\text{CPM{}}_{\text{IN}}, to mimic both the factual and counterfactual behavior of N\mathcal{N}{}. The key difference is how the information about the intervention Ci←c′{C_{i}}\leftarrow{c^{\prime}} is exposed to the model. Specifically, we adapt Interchange Intervention Training (Geiger et al., 2022) to train our CPMHI\text{CPM{}}_{\text{HI}} models for concept-based model explanation.

A conventional intervention on a hidden representation HH of a neural network N\mathcal{N}{} fixes the value of the representation HH to a constant. In an interchange intervention, we instead fix HH to the value it would have been when processing a separate source input ss. The result of the interchange intervention is a new model. Formally, we describe this new model as NH←Hs\mathcal{N}_{H\leftarrow H_{s}}, where ←\leftarrow is the conventional intervention operator and HsH_{s} is the value of hidden representation HH when processing input ss.

Here HCiH^{C_{i}} are hidden states designated for concept CiC_{i}. In essence, we train P\mathcal{P} to fully mediate the effect of intervening on CiC_{i} in the hidden representation HCiH^{C_{i}}. The source input ss is any input xu′,v′Ci=c′x_{u^{\prime},v^{\prime}}^{{C_{i}}={c^{\prime}}} that has Ci=c′{C_{i}}={c^{\prime}}. As P\mathcal{P} only receives information about the concept-level intervention Ci←c′{C_{i}}\leftarrow{c^{\prime}} via the interchange intervention HCi←HsCiH^{C_{i}}\leftarrow H_{s}^{C_{i}}, the model is forced to store all causally relevant information with regard to CiC_{i} in the corresponding hidden representation. This process is described in Figure 1(e).

At inference time, approximate counterfactuals are inaccessible, as before. To explain model N\mathcal{N}{} with regard to intervention Ci←c′{C_{i}}\leftarrow{c^{\prime}}, we manipulate the internal states of model P\mathcal{P} by intervening on the localized representation HCiH^{C_{i}} for concept CiC_{i}. To achieve this, we sample a source input xu′,v′Ci=c′x_{u^{\prime},v^{\prime}}^{{C_{i}}={c^{\prime}}} from the train set as any input xx that has Ci=c′{C_{i}}={c^{\prime}} to derive HsCiH_{s}^{C_{i}}.

Experiment Setup

CEBaB (Abraham et al., 2022) is a large benchmark of high-quality, labeled approximate counterfactuals for the task of sentiment analysis on restaurant reviews. The benchmark was created starting from a set of 2,299 original restaurant reviews from OpenTable. For each of these original reviews, approximate counterfactual examples were written by human annotators; the annotators were tasked to edit the original text to reflect a specific intervention, like ‘change the food evaluation from negative to positive’ or ‘change the service evaluation from positive to unknown’. In this way, the original reviews were expanded with approximate counterfactuals to a total of 15,089 texts. The groups of originals and corresponding approximate counterfactuals are partitioned over train, dev, and test sets. The pairs in the development and test set are used to benchmark explanation methods.

Each text in CEBaB was labeled by five crowdworkers with a 5-star sentiment score. In addition, each text was annotated at the concept level for four mediating concepts {Cambiance\{C_{\text{ambiance}}, CfoodC_{\text{food}}, CnoiseC_{\text{noise}}, and Cservice}C_{\text{service}}\}, using the labels {negative,unknown,positive}\{\text{negative},\text{unknown},\text{positive}\}, again with five crowdworkers annotating each concept-level label. We refer to Appendix A.1 and Abraham et al. 2022 for additional details.

As discussed above (Section 3 and Figure 1(b)), we consider two sources of approximate counterfactuals using CEBaB. For human-created counterfactuals, we use the edited restaurant reviews of the train set. For metadata-sampled counterfactuals, we sample factual inputs from the train set that have the desired combination of mediating concepts. Using all the human-created edits leads to 19,684 training pairs of factuals and corresponding approximate counterfactuals. Sampling counterfactuals leads to 74,574 pairs. We use these approximate counterfactuals to train explanation methods. Appendix A.2 provides more information about our pairing process.

2 Evaluation Metrics

This is simply the difference between the vectors of output scores for the two examples.

3 Baseline Methods

We compare our results with the best results obtained on the CEBaB benchmark. Crucially, BESTCEBaB\text{BEST}_{\text{CEBaB}} is not a single method, but rather pools together the best result obtained by any explanation method previously benchmarked on the CEBaB dataset, for every combination of model and metric.

S-Learner

Our version of S-Learner (Künzel et al., 2019) learns to mimic the factual behavior of black-box model N\mathcal{N}{} while making the intermediate concepts explicit.We use the finetuned concept-level sentiment analysis models released by Abraham et al. (2022). Given a factual input, a finetuned BERT model B\mathcal{B} first predicts values for the intermediate concepts. Then, a logistic regression model LRN\mathsf{LR}_{\mathcal{N}{}} is trained to map these intermediate concept values to the factual output of black-box model N\mathcal{N}{}, under the following objective.We use the default implementation LogisticRegression of scikit-learn (Buitinck et al., 2013).

By intervening on the intermediate predicted concept values at inference-time, we can hope to simulate the counterfactual behavior of N\mathcal{N}{}:

When using S-Learner in conjunction with approximate counterfactual inputs at train-time, we simply add this counterfactual data on top of the observational data that is typically used to train S-Learner.

GPT-3

Large language models such as GPT-3 (175B) have shown extraordinary power in terms of in-context learning (Brown et al., 2020).We use the largest davinci model publicly available at https://beta.openai.com/playground. We use GPT-3 to generate a new approximate counterfactual at inference time given a factual input and a descriptor of the intervention. This generated counterfactual is directly used to estimate the change in model behavior:

We use our train-time approximate counterfactual inputs to construct a prompt for GPT-3. Given this prompt, GPT-3 outputs new approximate counterfactuals given a factual and intervention descriptor. Full details on how these prompts are constructed can be found in Appendix A.7.

4 Causal Proxy Models

We train CPMs for the publicly available models released for CEBaB, fine-tuned as five-way sentiment classifiers on the factual data. This includes four model architectures: bert-base-uncased (BERT; Devlin et al. 2019), RoBERTa-base (RoBERTa; Liu et al. 2019), GPT-2 (GPT-2; Radford et al. 2019), and LSTM+GloVe (LSTM; Hochreiter & Schmidhuber 1997; Pennington et al. 2014). All Transformer-based models (Vaswani et al., 2017) have 12 Transformer layers. Before training, each CPM model is initialized with the architecture and weights of the black-box model we aim to explain. Thus, the CPMs are rooted in the factual behavior of N\mathcal{N}{} from the start. We include details about our setup in Appendix A.3.

The inference time comparisons for these models are as follows, where P\mathcal{P} in Eqn. 9 and Eqn. 10 refers to the CPM model trained under CPMIN\text{CPM{}}_{\text{IN}} and CPMHI\text{CPM{}}_{\text{HI}} objectives, respectively:

Here, ss is a source input with Ci=c′{C_{i}}={c^{\prime}}, and HCiH^{C_{i}} is the neural representation associated with CiC_{i} which takes value HsCiH_{s}^{C_{i}} on the source input ss. As HCiH^{C_{i}}, for BERT we use slices of width 192 taken from the 1st intermediate token of the 10th layer. For RoBERTa, we use the 8th layer instead. For GPT-2, we pick the final token of the 12th layer, again with slice width of 192. For LSTM, we consider slices of the attention-gated sentence embedding with width 64. Appendix A.5 studies the impact of intervention location and size.

Following the guidance on IIT given by Geiger et al. (2022), we train CPMHI\text{CPM{}}_{\text{HI}} with an additional multi-task objective as,

where our probe is parameterized by a multilayer perceptron MLP\mathsf{MLP}, and HxCiH_{x}^{C_{i}} is the value of hidden representation for the concept CiC_{i} when processing input xx with a concept label of cc for CiC_{i}.

Results

We first benchmark both CPM variants and our baseline methods on CEBaB. We show that the CPMs achieve state-of-the-art performance, for both types of approximate counterfactuals used during training (Section 5.1). Given the good factual performance achieved by CPMs, we subsequently investigate whether CPMs can be deployed both as predictor and explanation method at the same time (Section 5.2) and find that they can. Finally, we show that the localized representations of CPMHI\text{CPM{}}_{\text{HI}} give rise to concept-aware feature attributions (Section 5.3). Our supplementary materials report on detailed ablation studies and explore the potential of our methods for model debiasing.

Table 1 presents our main results. The results are grouped per approximate counterfactual type used during training. Both CPMIN\text{CPM{}}_{\text{IN}} and CPMHI\text{CPM{}}_{\text{HI}} beat BESTCEBaB{}_{\text{CEBaB}} in every evaluation setting by a large margin, establishing state-of-the-art explanation performance. Interestingly, CPMHI\text{CPM{}}_{\text{HI}} seems to slightly outperform CPMIN\text{CPM{}}_{\text{IN}} using sampled approximate counterfactuals, while slightly underperforming CPMIN\text{CPM{}}_{\text{IN}} on human-created approximate counterfactuals. Appendix A.6 reports on ablation studies that indicate that, for CPMHI\text{CPM{}}_{\text{HI}}, this state-of-the-art performance is primarily driven by the role of IIT in localizing concepts.

S-Learner, one of the best individual explainers from the original CEBaB paper (Abraham et al., 2022), shows only a marginal improvement when naively incorporating sampled and human-created counterfactuals during training over using no counterfactuals. This indicates that the large performance gains achieved by our CPMs over previous explainers are most likely due to the explicit use of a counterfactual training signal, and not primarily due to the addition of extra (counterfactual) data.

GPT-3 occasionally performs on-par with our CPMs, generally only slightly underperforming our best explainer on human-created counterfactuals, while being significantly worse on sampled counterfactuals. While the GPT-3 explainer also explicitly uses approximate counterfactual data, the results indicate that our proposed counterfactual mimic objectives give better results. The better performance of CPMs when considering sampled counterfactuals over GPT-3 shows that our approach is more robust to the quality of the approximate counterfactuals used. While the GPT-3 explainer is easy to set up (no training required), it might not be suitable for some explanation applications regardless of performance, due to the latency and cost involved in querying the GPT-3 API.

Across the board, explainers trained with human-created counterfactuals are better than those trained with sampled counterfactuals. This shows that the performance of explanation methods depends on the quality of the approximate counterfactual training data. While human counterfactuals give excellent performance, they may be expensive to create. Sampled counterfactuals are cheaper if the relevant metadata is available. Thus, under budgetary constraints, sampled counterfactuals may be more efficient.

Finally, CPMIN\text{CPM{}}_{\text{IN}} is conceptually the simpler of the two CPM variants. However, we discuss in Section 5.3 how the localized representations of CPMHI\text{CPM{}}_{\text{HI}} lead to additional explainability benefits.

2 Self-Explanation with CPM

As outlined in Section 3, CPMs learn to mimic both the factual and counterfactual behavior of the black-box models they are explaining. We show in Table 2 that our CPMs achieve a factual Macro-F1 score comparable to the black-box finetuned models.

We investigate if we can simply replace the black-box model with our CPM and use the CPM both as factual predictor and counterfactual explainer. To answer this questions, we measure the self-explanation performance of CPMs by simply replacing the black-box model N\mathcal{N}{} in Eqn. 5 with our factual CPM predictions at inference time.

Table 3 reports these results. We find that both CPMIN\text{CPM{}}_{\text{IN}} and CPMHI\text{CPM{}}_{\text{HI}} achieve better self-explanation performance compared to providing explanations for another black-box model. Furthermore, CPMHI\text{CPM{}}_{\text{HI}} provides better self-explanation than CPMIN\text{CPM{}}_{\text{IN}}, suggesting our interchange intervention procedure leads the model to localize concept-based information in hidden representations. This shows that CPMs may be viable as replacements for their black-box counterpart, since they provide similar task performance while providing faithful counterfactual explanations of both the black-box model and themselves.

We have shown that CPMHI\text{CPM{}}_{\text{HI}} provides trustworthy explanations (Section 5.1). We now investigate whether CPMHI\text{CPM{}}_{\text{HI}} learns representations that mediate the effects of different concepts. We adapt Integrated Gradients (IG; Sundararajan et al. 2017) to provide concept-aware feature attributions, by only considering gradients flowing through the hidden representation associated with a given concept. We formalize this version of IG in Appendix A.8.

In Table 4, we compare concept-aware feature attibutions for two variants of CPMHI\text{CPM{}}_{\text{HI}} (IIT and Multi-task) and the original black-box (Finetuned) model. For IIT we remove the multi-task objective LMulti\mathcal{L}_{\text{Multi}} during training and for Multi-task we remove the the interchange intervention objective LHI\mathcal{L}_{\text{HI}}. This helps isolate the individual effects of both losses on concept localization. All three models predict a neutral final sentiment score for the considered input, but they show vastly different feature attributions. Only IIT reliably highlights words that are semantically related to each concept. For instance, when we restrict the gradients to flow only through the intervention site of the noise concept, “loud” is the word highlighted the most that contributes negatively. When we consider the service concept, words like “friendly” and “waiter” are highlighted the most as contributing positively. These contrasts are missing for representations of the Multi-task and Finetuned models. Only the IIT training paradigm pushes the model to learn causally localized representations. For the service concept, we notice that the IIT model wrongfully attributes “delicious”. This could be useful for debugging purposes and could be used to highlight potential failure modes of the model.

Conclusion

We explored the use of approximate counterfactual training data to build more robust causal explanation methods. We introduced Causal Proxy Models (CPMs), which learn to mimic both the factual and counterfactual behaviors of a black-box model N\mathcal{N}{}. Using CEBaB, a benchmark for causal concept-based explanation methods, we demonstrated that both versions of our technique (CPMIN\text{CPM{}}_{\text{IN}} and CPMHI\text{CPM{}}_{\text{HI}}) significantly outperform previous explanation methods.

Interestingly, we find that our GPT-3 based explanation method performs on-par with our best CPM model in some settings. While test-time use of GPT-3 as explanation method might not be feasible, we believe this result shows that GPT-3 could be deployed to supplement human-annotation efforts for counterfactual data creation.

Our results suggest that CPMs can be more than just explanation methods. They achieve factual performance on par with the model they aim to explain, and they can explain their own behavior. This paves the way to using them as deployed models that both perform tasks and offer explanations. In addition, the causally localized representations of our CPMHI\text{CPM{}}_{\text{HI}} variant are very intuitive, as revealed by our concept-aware feature attribution technique. We believe that causal localization techniques could play a vital role in further model explanation efforts.

Acknowledgement

This research is supported in part by a grant from Meta AI. Karel D’Oosterlinck was supported through a doctoral fellowship from the Special Research Fund (BOF) of Ghent University.

References

Appendix A Appendix

Table 5 shows dataset statistics of CEBaB. The variants of CEBaB we consider only impact the train split. The top panel shows the number of observational samples and edits introduced in the CEBaB paper. The bottom panel shows our paired versions, where we create approximate counterfactual pairs. We explore two variants of approximate counterfactuals: human-created and sampled counterfactuals (Section 4.1). The human setting considers all pairs made possible by using all data. The sampling setting considers pairs sampled from only the observational data, as discussed in Section A.2.

A.2 Types of Approximate Counterfactual Pairs

Our approximate counterfactual training data comes in paired sentences of (original sentence, approximate counterfactual sentence). The approximate counterfactuals differs from their original counterparts in only one concept value. We consider approximate counterfactual pairs to be symmetric: we use both (original sentence, approximate counterfactual sentence) and (approximate counterfactual sentence, original sentence) as training pairs.

CEBaB contains multiple counterfactual sentences for each original review. To achieve this, the dataset creators asked annotators to edit the original sentence to achieve a specified goal (e.g., ‘change the evaluation of the restaurant’s food to negative’). These originals and corresponding edits form our human pairs.

Metadata-sampled Counterfactuals

Human-created counterfactuals are not always available. With CEBaB, we simulate a second type of approximate counterfactuals by using metadata-guided heuristics: for a given original sentence, we sample a counterfactual from the train set by matching concept labels while allowing only one label to be changed.

During training, we also consider null effect pairs in our sampling setup. These pairs resemble cases where our approximate counterfactual sentence is identical to the original sentence. When training our models on these pairs, we expect our models to predict the same counterfactual and factual output.

A.3 Training Regimes

To train CPMIN\text{CPM{}}_{\text{IN}}, we use the same model architecture as N\mathcal{N}, and initialize it with the model weights using weights from N\mathcal{N}. The maximum number of training epochs is set to 30 with a learning rate of 5e−55e^{-5} and an effective batch size of 128. The learning rate linearly decays to over the 30 training epochs. We employ an early stopping strategy for COSICaCE\text{COS}_{\texttt{ICaCE}} over the dev set for an interval of 50 steps with early stopping patience set to 20. We set the max sequence length to 128 and the dropout rate to 0.10.1. We take a weighted sum of two objectives as the loss term for training CPMHI\text{CPM{}}_{\text{HI}}. Specifically, we use [wMimic,wIN]=[1.0,3.0][w_{\text{Mimic}},w_{\text{IN}}]=[1.0,3.0]. For the smoothed cross-entropy loss, we use a temperature of 2.02.0.

To train CPMHI\text{CPM{}}_{\text{HI}}, we use the same model architecture as N\mathcal{N}, and initialize it with the model weights using weights from N\mathcal{N}. The maximum number of training epochs is set to 30 with a learning rate of 8e−58e^{-5} and an effective batch size of 256. We use a higher learning rate of 0.0010.001 for the LSTM model as it enables quicker convergence. The learning rate linearly decays to over the 30 training epochs. We employ an early stopping strategy for COSICaCE\text{COS}_{\texttt{ICaCE}} over the dev set for an interval of 10 steps with early stopping patience set to 20. We set the max sequence length to 128 and the dropout rate to 0.10.1. We take a weighted sum of three objectives as the loss term for training CPMHI\text{CPM{}}_{\text{HI}}. Specifically, we use [wMimic,wMulti,wHI]=[1.0,1.0,3.0][w_{\text{Mimic}},w_{\text{Multi}},w_{\text{HI}}]=[1.0,1.0,3.0]. In Appendix A.6, we conduct a set of ablation studies to isolate the individual contributions from each objective. For the smoothed cross-entropy loss, we use a temperature of 2.02.0.

Our models are all implemented in PyTorch (Paszke et al., 2019) and using the HuggingFace library (Wolf et al., 2019). All of our results are aggregated over three distinct random seeds. To foster reproducibility, we will release our code repository and model artifacts to the public.

A.4 Additional Baseline Results

Table 6 shows baselines adapted from Abraham et al. (2022), which contains the present state-of-the-art explanation methods for the CEBaB benchmark. We report the best scores across these explanation methods in Table 1. These baselines are trained without using counterfactual data. Thus, we build additional baselines that use counterfactual data as shown in Table 7. S-Learner is selected as the best performing models and included in Table 1 for comparisons. The equations for the additional baselines are as follows:

where srandoms^{\text{random}} is a randomly sampled training input, sapproxs^{\text{approx}} is a training input sampled to match the concept-level labels of the true counterfactual under intervention Ci←c′{C_{i}}\leftarrow{c^{\prime}}, DCi←c′\mathcal{D}^{{C_{i}}\leftarrow{c^{\prime}}} is the set of all approximate counterfactual training pairs that represent a Ci←c′{C_{i}}\leftarrow{c^{\prime}} intervention, and ff is a look-up function that returns the ground-truth label associated with an input.

The signatures of EATE\mathcal{E}^{\text{ATE}} and ENCaCE\mathcal{E}_{\mathcal{N}}^{\text{CaCE}} reflect that they are independent of the specific factual input xu,vx_{u,v} considered. Furthermore, EATE\mathcal{E}^{\text{ATE}} is independent of N\mathcal{N}{} given that this explainer only uses ground-truth training labels to estimate causal effects.

A.5 Intervention Site Location and Size

Previous work shows that neurons in different layers and groups can encode different high-level concepts (Vig et al., 2020; Koh et al., 2020). CPMHI\text{CPM{}}_{\text{HI}} pushes concept-related information to localize at the targeted intervention site (the aligned neural representations for each concept). In this section, we investigate how the location and the size of the intervention site impact CPMHI\text{CPM{}}_{\text{HI}} performance. We use the optimal location and size found in this study for other results presented in this paper.

For Transformer-based models, we vary the location of the intervention site by intervening on the “[CLS]” token embedding layer ll. Specifically, we set l={2,4,6,8,10,12}l=\{2,4,6,8,10,12\}. We skip this experiment for non-Transformer-based model (i.e., LSTM) since it only contains a single sentence embedding.

As shown in the top panel of Figure 2, intervention location significantly affects CPMHI\text{CPM{}}_{\text{HI}} performance. Our results show that layer 10 for BERT, layer 8 for RoBERTa, and layer 12 for GPT-2 lead to the best performance. This suggests layers have different efficacy in terms of information localization. Our results also show that intervening with deeper layers tends to provide better performance. However, for both BERT and RoBERTa, intervening on the last layer results in a slightly worse performance compared to earlier layers. This suggests that leaving Transformer blocks after the intervention site helps localized information to be processed by the neural network.

Size

For Transformer-based models, we change the size of the intervention site dcd_{c} for each concept. Specifically, we set dc={1,16,64,128,192}d_{c}=\{1,16,64,128,192\}. For instance when dc=1d_{c}=1, we use a single dimension of the “[CLS]” token embedding to represent each concept, starting from the first dimension of the vector. For our non-Transformer-based model (LSTM), we intervene on the attention-gated sentence embedding whose dimension size is set to 300. Accordingly, we set dc={1,16,64,75}d_{c}=\{1,16,64,75\}.

As shown in Figure 2, larger intervention sites lead to better performance for all Transformer-based models. For LSTM, we find that the optimal size is the second largest one instead. On the other hand, our results suggest that the performance gain from the increase of size diminishes as we increase the size for all model architectures.

Geiger et al. (2022) show that training with a multi-task objective helps IIT to improve generalizability. In this experiment, we aim to investigate whether the multi-task objective we added for CPMHI\text{CPM{}}_{\text{HI}} plays an important role in achieving good performance. Specifically, we conduct two ablation studies: removing the multi-task objective by setting wMulti=0.0w_{\text{Multi}}=0.0, and removing the IIT objective by setting wHI=0.0w_{\text{HI}}=0.0.

Table 8 shows our results, which demonstrate that the IIT objective is the main factor that drives CPMHI\text{CPM{}}_{\text{HI}} performance. Our results also suggest that the multi-task objective brings relatively small but consistent performance gains. Overall, our findings corroborate those of Geiger et al. (2022) and provide concrete evidence that the combination of two objectives always results in the best-performing explanation methods across all model architectures.

Additionally, we explore two baselines for CPMHI\text{CPM{}}_{\text{HI}}. Firstly, we randomly initialize the weights of CPMHI\text{CPM{}}_{\text{HI}}. Secondly, we take the original black-box model as our CPMHI\text{CPM{}}_{\text{HI}}. Compared to the results in Table 1, these two baselines fail catastrophically, suggesting the importance of our IIT paradigm.

As mentioned in Section 3, we sample a source input xu′,v′Ci=c′x_{u^{\prime},v^{\prime}}^{{C_{i}}={c^{\prime}}} from the train set as any input xx that has Ci=c′{C_{i}}={c^{\prime}} to estimate the counterfactual output. Furthermore, we explore two additional sampling strategies. First, we create a baseline where we randomly sample a source input from the train without any concept label matching. Second, we sample a source input from the train set using the predicted concept label of our multi-task probe, instead of the true concept label from the dataset.

As shown in Table 9, the quality of our source inputs impact our performance significantly. For instance, when sampling source input at random, CPMHI\text{CPM{}}_{\text{HI}} fails catastrophically for all evaluation metrics. On the other hand, when we sampling source based on the predicted labels using the multi-task probe, CPMHI\text{CPM{}}_{\text{HI}} maintains its performance.

A.7 GPT-3 Generation Process

For each few-shot learning prompt, we insert an initial string of the form of “Make the following restaurant reviews include c′c^{\prime} mentions of CiC_{i}.”, where c′c^{\prime} is expressed as one of {“POSITIVE”, “NEGATIVE”, “NOT” } (“NOT” corresponds to making the review be unknown regarding the concept CiC_{i}) and CiC_{i} is one of {“AMBIANCE”, “FOOD”, “NOISE”, “SERVICE”}. We sample using a temperature of 0.9, without any frequency or presence penalties (since we expect the counterfactual review to be similar to the original review). In preliminary experimentation, we found that capitalizing the mediating concept and target value results and inserting line breaks between examples made for better completions, although there is room for future research in this area.

We used the OpenAI API to access GPT-3. At the current price rate of 0.02per1,000tokens,thetotalcostofcreatingourcounterfactuals(around4,000examples)wasapproximately0.02 per 1,000 tokens, the total cost of creating our counterfactuals (around 4,000 examples) was approximately50 per approximate counterfactuals creation strategy.

A.8 Integrated Gradients

We adapt the Integrated Gradients (IG) method of Sundararajan et al. (2017) to qualitatively assess whether CPMHI\text{CPM{}}_{\text{HI}} learned explainable representations of mediated concepts at its intervention sites. The IG algorithm computes the average gradient from the model output to its input by incrementally interpolating from a “blank” input x′x^{\prime} (consisting only of “[PAD]” tokens) to the original input xx. Eqn. 16 is the integrated gradients equation originally proposed in Sundararajan et al. (2017), applied to a CPM model P\mathcal{P} on input xx.

Here, ∂P(x)∂xj\frac{\partial\mathcal{P}(x)}{\partial x_{j}} is the derivative of P\mathcal{P} on the jjth dimension of xx.

In our implementation of IG, we wish to show the per-token attribution of input xx on the model’s final output P(x)\mathcal{P}(x), mediated by the hidden representation of a concept in P\mathcal{P}. That is, we’d like to ask, “What is the effect of the word ‘delicious’ in the input on the model’s output, when we restrict our focus only on the model’s representation of the concept food?”

To answer this question, we compute the gradient of the model output P(x)\mathcal{P}(x) with respect to the input xx but restrict the gradient to flow through the intervention site for a particular concept. This allows us to capture the per-token attribution of the model’s final output (whether particular words contributed to a positive, negative, or neutral sentiment prediction), mediated by the concept that is represented by the specified intervention site. For example, in Table 4, we can see that “delicious” has a positive attribution to the output of the model when we focus on its representation of the concept food.

Formally, consider a trained CPM model P\mathcal{P}, an input xx and mediating concept CiC_{i}. Let HCiH^{C_{i}} be the activation of P\mathcal{P} at the intervention site for CiC_{i}. We define the gradient of P(x)\mathcal{P}(x) along dimension jj, mediated by CiC_{i}, as

Eqn. 17 restricts the gradient to only flow through the hidden representation of the concept along which we’d like to interpret our model.

We integrate these mediated gradients over a straight path between input xx and baseline x′x^{\prime}, analogous to Eqn. 16. We implement our IG method using CaptumAI library.https://captum.ai/ We use the default parameters for our runs with number of iterations set to 50, and we set the integral method as gausslegendre. We set the multiply-by-inputs flag to True. To visualize individual word importance, we conduct zz-score normalization of attribution scores over input tokens per each concept, and then linearly scale scores between [−1-1, +1+1].

Table 10 extends Table 4 in our main text with additional ablation studies on our training objectives.

A.9 Model Debiasing

Being able to accurately predict outputs for counterfactual inputs enables explanation methods to faithfully debias a model with regard to a desired concept. For instance, with CEBaB, debiasing a concept (e.g., “food”) is equivalent to estimating the counterfactual output when we set the concept label for a concept to be unknown.

In this section, we briefly study the extent to which the CPMHI\text{CPM{}}_{\text{HI}} can function as a debiasing method. To debias a concept, we enforce the sampled source input ss as in Eqn. 3 to have unknown as its concept label for the concept to be debiased.

To show our methods can faithfully debias a targeted concept, we evaluate the correlations between the predicted overall sentiment label for sentences and the concept labels for each concept. Without any debiasing technique, we expect concept labels to be highly correlated with the overall sentiment label (e.g., if food is positive, it is more likely that the overall sentiment is positive). We use CPMHI\text{CPM{}}_{\text{HI}} trained for the BERT model architecture as an example, and use examples in the test set.

Figure 5 shows correlation plots for the original Finetuned model as well as CPMHI\text{CPM{}}_{\text{HI}}. As expected, the correlation of the food concept is weakened through the debiasing pipeline by 57.50%. Our results also suggest that correlations of other concepts are affected, which suggests a future research direction focused on minimizing the impact of the debiasing pipeline on irrelevant concepts. We include results for the remaining concepts in the Appendix A.9.

Figure 5(a) to Figure 5(d) show debiasing visualizations for three concepts: ambiance, noise and service. We use a CPMHI\text{CPM{}}_{\text{HI}} for the BERT model architecture as an example. We calculate the distributions with examples in the test set.

A.10 Learning Dynamics

Figure 6 shows three different metrics measured on the dev and the test sets for a CPMHI\text{CPM{}}_{\text{HI}} trained for the BERT model architecture as an example. Since we use COSICaCE\text{COS}_{\texttt{ICaCE}} on the dev set to early stop our training process, we find our CPMHI\text{CPM{}}_{\text{HI}} reaches a local minimum on COSICaCE\text{COS}_{\texttt{ICaCE}} while L2ICaCE\text{L2}_{\texttt{ICaCE}} and NormDiffICaCE\text{NormDiff}_{\texttt{ICaCE}} are still trending downward. This suggests future research may need to choose desired metrics to optimize for during training, for early stopping to reach the best performing model.

Table 11 visualizations of word importance scores using our version of Integrated Gradient (IG). Different from Table 4 and Table 10, which show the visualizations of our optimized model, we show a per-epoch result for for CPMHI\text{CPM{}}_{\text{HI}}, followed with our best model appended at the end. Our results suggest that early checkpoints in the training process focus at drastically different input words comparing to later checkpoints, though all models predict neutral for this given sentence. In addition, gradient aggregations over input words are rather stable towards the end the training. More importantly, CPMHI\text{CPM{}}_{\text{HI}} learns how to highlight words that are semantically related to each concept gradually. For instance, we can see a clear trend of emphasising the word “decorations” for the ambiance concept throughout the training process. This suggests that our training procedure induces causally motivated gradients over input words gradually through the training process.