Finding Alignments Between Interpretable Causal Variables and Distributed Neural Representations
Atticus Geiger, Zhengxuan Wu, Christopher Potts, Thomas Icard, Noah D. Goodman
Introduction
Can an interpretable symbolic algorithm be used to faithfully explain a complex neural network model? This is a key question for interpretability; a positive answer can provide guarantees about how the model will behave, and a negative answer could lead to fundamental concerns about whether the model will be safe and trustworthy.
Causal abstraction provides a mathematical framework for precisely characterizing what it means for any complex causal system (e.g., a deep learning model) to implement a simpler causal system (e.g., a symbolic algorithm) . For modern AI models, the fundamental operation for assessing whether this relationship holds in practice has been the interchange intervention (also known as activation patching), in which a neural network is provided a ‘base’ input, and sets of neurons are forced to take on the values they would have if different ‘source’ inputs were processed . The counterfactuals that these interventions create are the basis for causal inferences about model behavior.
Geiger et al. show that the relevant causal abstraction relation obtains when interchange interventions on aligned high-level variables and low-level variables have equivalent effects. This ideal relationship rarely obtains in practice, but the proportion of interchange interventions with the same effect (interchange intervention accuracy; IIA) provides a graded notion, and Geiger et al. formally ground this metric in the theory of approximate causal abstraction. Geiger et al. also use causal abstraction theory as a unified framework for a wide range of recent intervention-based analysis methods .
Causal abstraction techniques have been applied to diverse problems . However, previous applications have faced two central challenges. First, causal abstraction requires a computationally intensive brute-force search process to find optimal alignments between the variables in the high-level model and the states of the low-level one. Where exhaustive search is intractable, we risk missing the best alignment entirely. Second, these prior methods are localist: they artificially limit the space of possible alignments by presupposing that high-level causal variables will be aligned with disjoint groups of neurons. There is no reason to assume this a priori, and indeed much recent work in model explanation (see especially Ravfogel et al. 30, 31, Elazar et al. 8, Olah et al. 26, Olsson et al. 27) is converging on the insight of Smolensky , Rumelhart et al. , and McClelland et al. that individual neurons can play multiple conceptual roles. Smolensky identified distributed neural representations as “patterns” consisting of linear combinations of unit vectors.
In the current paper, we propose distributed alignment search (DAS), which overcomes the above limitations of prior causal abstraction work. In DAS, we find the best alignment via gradient descent rather than conducting a brute-force search. In addition, we use distributed interchange interventions, which are “soft” interventions in which the causal mechanisms of a group of neurons are edited such that (1) their values are rotated with a change-of-basis matrix, (2) the targeted dimensions of the rotated neural representation are fixed to be the corresponding values in the rotated neural representation created for the source inputs, and (3) the representation is rotated back to the standard neuron-aligned basis. The key insight is that viewing a neural representation through an alternative basis that is not aligned with individual neurons can reveal interpretable dimensions .
In our experiments, we evaluate the capabilities of DAS to provide faithful and interpretable explanations with two tasks that have obvious interpretable high-level algorithmic solutions with two intermediate variables. In both tasks, the distributed alignment learned by DAS is as good or better than both the closest localist alignment and the best localist alignment in a brute-force search.
In our first set of experiments, we focus on a hierarchical equality task that has been used extensively in developmental and cognitive psychology as a test of relational reasoning : the inputs are sequences , and the label is given by . We train a simple feed-forward neural network on this task and show that it perfectly solves the task. Our key question: does this model implement a program that computes and as intermediate values, as we might hypothesize humans do? Using DAS, we find a distributed alignment with 100% IIA. In other words, the network is perfectly abstracted by the high-level model; the distinction between the learned neural model and the symbolic algorithm is thus one of implementation.
Our second task models a natural language inference dataset where the inputs are premise and hypothesis sentences that are identical but for the words and ; the label is either entails ( makes true) or contradicts/neutral ( makes false). We fine-tune a pretrained language model to perfectly solve the task. With DAS, we find a perfect alignment (100% IIA) to a causal model with a binary variable for the entailment relation between the words and (e.g., dog entails mammal).
In both our sets of experiments, the DAS analyses reveal perfect abstraction relations. However, we also identify an important difference between them. In the NLI case, the entailment relation can be decomposed into representations of and . What appears to be a representation of lexical entailment is, in this case, a “data structure” containing two representations of word identity, rather than an encoding of their entailment relation. By contrast, the hierarchical equality models learn representations of and that cannot be decomposed into representations of , , and . In other words, these relations are entirely abstracted from the entities participating in the relation; DAS reveals that the neural network truly implements a symbolic, tree-structured algorithm.
Related Work
A theory of causal abstraction specifies exactly when a ‘high-level causal model’ can be seen as an abstract characterization of some ‘low-level causal model’ . The basic idea is that high-level variables are associated with (potentially overlapping) sets of low-level variables that summarize their causal mechanisms with respect to a set of hard or soft interventions . In practice, a graded notion of approximate causal abstraction is often more useful .
Geiger et al. argue that causal abstraction is a generic theoretical framework for providing faithful and interpretable explanations of AI models and show that LIME , causal effect estimation , causal mediation analysis , iterated nullspace projection , and circuit-based explanations can all be understood as causal abstraction analysis.
Interchange intervention training (IIT) objectives are minimized when a high-level causal model is an abstraction of a neural network under a given alignment . In this paper, we use IIT objectives to learn an alignment between a high-level causal model and a deep learning model.
Methods
We focus on acyclic causal models and seek to provide an intuitive overview of our method. An acyclic causal model consists of input, intermediate, and output variables, where each variable has an associated set of values it can take on and a causal mechanism that determine the value of the variable based on the value of its causal parents. For a simple running example, we modify the boolean conjunction models of to reveal key properties of DAS. A causal model for this problem can be defined as below, where the inputs and outputs are booleans t and f. Alongside , we also define a causal model of a linear feed-forward neural network that solves the task. Here we show , , and the parameters of :
The model predicts t if and f otherwise. This network solves the boolean conjunction problem perfectly in that all pairs of input boolean values are mapped to the intended output.
An input of a model determines a unique total setting of all the variables in the model. The inputs are fixed to be and the causal mechanisms of the model determine the values of the remaining variables. We denote the values that assigns to the variable or variables as . For example, .
Interventions are a fundamental building block of causal models, and of causal abstraction analysis in particular. An intervention is a setting of variables . Together, an intervention and an input setting of a model determine a unique total setting that we denote as . The inputs are fixed to be , and the causal mechanisms of the model determine the values of the non-intervened variables, with the intervened variables being fixed to .
We can define interventions on both our causal model and our neural model . For example, is our boolean model when it processes input but with variable set to t. This has the effect of changing the output value to t. Similarly, whereas leads to an intermediate values and and output value , if we compute , then the output value is . This has the effect of changing the predicted value to t, because .
2 Alignment
In causal abstraction analysis, we ask whether a specific low-level model like implements a high-level algorithm like . This is always relative to a specific alignment of variables between the two models. An alignment assigns to each high-level variable a set of low-level variables and a function that maps from values of the low-level variables in to values of the aligned high-level variable . One possible alignment between and is shown in the diagram above: is depicted by the dashed lines connecting and .
We immediately know what the functions for high-level input and output variables are. For the inputs, t is encoded as and f is encoded as , meaning and . For the output, the network only predicts t if , meaning if , else f. This is simply a consequence of how a neural network is used and trained. The functions for high-level intermediate variables and must be discovered and verified experimentally.
3 Constructive Causal Abstraction
Relative to an alignment like this, we can define abstraction:
(Constructive Causal Abstraction) A high-level causal model is a constructive abstraction of a low-level causal model under alignment exactly when the following holds for every low-level input setting and low-level intervention :
being a causal abstraction of under guarantees that the causal mechanism for each high-level variable is a faithful rendering of the causal mechanisms for the low-level variables in .
To assess the degree to which a high-level model is a constructive causal abstraction of a low-level model, we perform interchange interventions:
(Interchange Interventions) Given source input settings , and non-overlapping sets of intermediate variables for model , define the interchange intervention as the model
where concatenates a set of interventions.
A base input setting can be fed into the resulting model to compute the counterfactual output value. Consider the following interchange intervention:
We process a base input and a source input, and then we intervene on a target variable, replacing it with the value obtained by processing the source. Our causal model is fully known, and so we know ahead of time that this interchange intervention yields t. For our neural network, the corresponding behavior is not known ahead of time. The interchange intervention corresponding to the above (according to the alignment we are exploring) is as follows
And, indeed, the counterfactual behavior of the model and the network are unequal:
ftV_{1}=\textsc{t}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi>V</mi><mn>2</mn></msub><mo>=</mo><mstyle mathcolor="#cc0000"><mtext>\textsc</mtext></mstyle><mi>t</mi></mrow><annotation encoding="application/x-tex">V_{2}=\textsc{t}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8333em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.2222em;">V</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:-0.2222em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">2</span></span></span></span></span><span class="vlist-s"></span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="base"><span class="strut" style="height:1em;vertical-align:-0.25em;"></span><span class="mord text" style="color:#cc0000;"><span class="mord" style="color:#cc0000;">\textsc</span></span><span class="mord"><span class="mord mathnormal">t</span></span></span></span></span></span>V_{3}=\textsc{t}ttV_{1}=\textsc{t}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi>V</mi><mn>2</mn></msub><mo>=</mo><mstyle mathcolor="#cc0000"><mtext>\textsc</mtext></mstyle><mi>t</mi></mrow><annotation encoding="application/x-tex">V_{2}=\textsc{t}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8333em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.2222em;">V</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:-0.2222em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">2</span></span></span></span></span><span class="vlist-s"></span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="base"><span class="strut" style="height:1em;vertical-align:-0.25em;"></span><span class="mord text" style="color:#cc0000;"><span class="mord" style="color:#cc0000;">\textsc</span></span><span class="mord"><span class="mord mathnormal">t</span></span></span></span></span></span>V_{3}=\textsc{t} f Under the given alignment, the interchange interventions at the low and high level have different effects. Thus, we have a counterexample to constructive abstraction as given in Definition 3.1. Although has perfect behavioral accuracy, its accuracy under the counterfactuals created by our interventions is not perfect, and thus is not a constructive abstraction of under this alignment.
4 Distributed Interventions
The above conclusion is based on the kind of localist causal abstraction explored in the literature to date. As noted in Section 1, there are two risks associated with this conclusion: (1) we may have chosen a suboptimal alignment, and (2) we may be wrong to assume that the relevant structure will be encoded in the standard basis we have implicitly assumed throughout.
If we simply rotate the representation by to get a new representation , then the resulting network has perfect behavioral and counterfactual accuracy when we align and with and . What this reveals is that there is an alignment, but not in the basis we chose. Since the choice of basis was arbitrary, our negative conclusion about the causal abstraction relation was spurious.
This rotation localizes the information about the first and second argument into separate dimensions. To understand this, observe that the weight matrix of the linear network rotates a two dimensional vector by and the rotation matrix rotates the representation by . The two matrices are inverses. Because this network is linear, there is no activation function and so rotating the hidden representation “undoes” the transformation of the input by the weight matrix. Under this non-standard basis, the first hidden dimension is equal to the first input argument and the second hidden dimension is equal to the second input argument.
This reveals an essential aspect of distributed neural representations: there is a many-to-many mapping between neurons and concepts, and thus multiple high-level causal variables might be encoded in structures from overlapping groups of neurons . In particular, Smolensky proposes that viewing a neural representation under a basis that is not aligned with individual neurons can reveal the interpretable distributed structure of the neural representations.
To make good on this intuition we define a distributed intervention, which first transforms a set of variables to a vector space, then does interchange on orthogonal sub-spaces, before transforming back to the original representation space.
Distributed Interchange Interventions We begin with a causal model with input variables and source input settings . Let be a subset of variables in , the target variables. Let be a vector space with subspaces that form an orthogonal decomposition, i.e., . Let be an invertible function . Write for the orthogonal projection operator of a vector in onto subspace .Thus, generalizes to arbitrary vector spaces. A distributed interchange intervention yields a new model which is identical to except that the mechanisms (which yield values of from a total setting) are replaced by:
Notice that in this definition the base setting is partially preserved through the intervention (in subspace ) and hence this is a soft intervention on that rewrites causal mechanisms while maintaining a causal dependence between parent and child.
Under this new alignment, the high-level interchange intervention is aligned with the low-level distributed interchange intervention
and the counterfactual output behavior of and are equal:
1<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mn>1</mn></mrow><annotation encoding="application/x-tex">1</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6444em;"></span><span class="mord">1</span></span></span></span></span>H_{1}=0.6<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi>H</mi><mn>2</mn></msub><mo>=</mo><mn>1.28</mn></mrow><annotation encoding="application/x-tex">H_{2}=1.28</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8333em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0813em;">H</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:-0.0813em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">2</span></span></span></span></span><span class="vlist-s"></span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="base"><span class="strut" style="height:0.6444em;"></span><span class="mord">1.28</span></span></span></span></span>O=0.08<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mn>1</mn></mrow><annotation encoding="application/x-tex">1</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6444em;"></span><span class="mord">1</span></span></span></span></span>H_{1}=-0.34<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi>H</mi><mn>2</mn></msub><mo>=</mo><mn>0.94</mn></mrow><annotation encoding="application/x-tex">H_{2}=0.94</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8333em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0813em;">H</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:-0.0813em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">2</span></span></span></span></span><span class="vlist-s"></span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="base"><span class="strut" style="height:0.6444em;"></span><span class="mord">0.94</span></span></span></span></span>1.0<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mn>1.0</mn></mrow><annotation encoding="application/x-tex">1.0</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6444em;"></span><span class="mord">1.0</span></span></span></span></span>0.0<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mn>1.0</mn></mrow><annotation encoding="application/x-tex">1.0</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6444em;"></span><span class="mord">1.0</span></span></span></span></span>\Bigg{[}\begin{array}[]{rr}\cos(-20^{\circ})&-\sin(-20^{\circ})\\ \sin(-20^{\circ})&\phantom{-}\cos(-20^{\circ})\end{array}\Bigg{]}<span class="katex-error" title="ParseError: KaTeX parse error: Invalid delimiter type 'ordgroup' at position 6: \Bigg{̲[̲}̲\begin{array}[]…" style="color:#cc0000">\Bigg{[}\begin{array}[]{rr}\cos(-20^{\circ})&-\sin(-20^{\circ})\\ \sin(-20^{\circ})&\phantom{-}\cos(-20^{\circ})\end{array}\Bigg{]}</span>\Bigg{[}\begin{array}[]{rr}\cos(20^{\circ})&-\sin(20^{\circ})\\ \sin(20^{\circ})&\phantom{-}\cos(20^{\circ})\end{array}\Bigg{]}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi>H</mi><mn>1</mn></msub><mo>=</mo><mn>0.6</mn></mrow><annotation encoding="application/x-tex">H_{1}=0.6</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8333em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0813em;">H</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:-0.0813em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">1</span></span></span></span></span><span class="vlist-s"></span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="base"><span class="strut" style="height:0.6444em;"></span><span class="mord">0.6</span></span></span></span></span>H_{2}=1.28t In what follows we will assume that are already vector spaces (which is true for neural nets) and the functions are rotation operators. In this case, the subspaces can be identified without loss of generality with those spanned by the first basis vectors for , the next basis vectors for , and so on. (The following methods would be well-defined for non-linear transformations, as long as they were invertible and differentiable, but efficient implementation becomes harder.)
5 Distributed Alignment Search
The question then arises of how to find good rotations. As we discussed above, previous causal abstraction analyses of neural networks have performed brute-force search through a discrete space of hand-picked alignments. In distributed alignment search (DAS), we find an alignment between one or more high-level variables and disjoint sub-spaces (but not necessarily subsets) of a large neural representation. We define a distributed interchange intervention training objective, use differentiable parameterizations for the space of orthogonal matrices (such as provided by PyTorch), and then optimize the objective with stochastic gradient descent. Crucially, the low-level and high-level models are frozen during learning so we are only changing the alignment.
In the following definition we assume that a neural network specifies an output distribution for a given input, which can then be pushed forward to a distribution on output values of the high-level model via an alignment function . We may similarly interpret even a deterministic high-level model as defining a (e.g., delta) distribution on output values. We make use of these distributions, after interchange intervention, to define a differentiable loss for the rotation matrix which aligns intermediate variables.
Distributed Interchange Intervention Training Objective Begin with a low-level neural network , with low-level input settings , a high-level algorithm , with high-level output settings , and an alignment for their input and output variables. Suppose we want to align intermediate high level variables with rotated subspaces of a neural representation with learned rotation matrix .
In general, we can define a training objective using any differentiable loss function that quantifies the distance between two total high-level settings.
While we still have discrete hyperparameters —the target population and the dimensionality of the sub-spaces used for each high-level variable—we may use stochastic gradient descent to determine the rotation that minimizes loss, thus yielding the best distributed alignment between and .
6 Approximate Causal Abstraction
Perfect causal abstraction relationships are unlikely to arise for neural networks trained to solve complex empirical tasks. We use a graded notion of accuracy:
Distributed Interchange Intervention Accuracy Given low-level and high-level causal models and with alignment , rotation , and orthogonal decomposition . If we let be low-level input settings and be high-level intermediate variables the interchange intervention accuracy (IIA) is as follows
IIA is the proportion of aligned interchange interventions that have equivalent high-level and low-level effects. In our example with and , IIA is 100% and the high-level model is a perfect abstraction of the low-level model (Def. 3.1). When IIA is 100%, we rely on the graded notion of -on-average approximate causal abstraction , which coincides with IIA.
7 General Experimental Setup
We illustrate the value of DAS by analyzing feed-forward networks trained on a hierarchical equality and pretrained Transformer-based language models fine-tuned on a natural language inference task. Our evaluation paradigm is as follows:
Train the neural network to solve the task. In all experiments, the neural models achieve perfect accuracy on both training and testing data.
Create interchange intervention training datasets using a high-level causal model. Each example consists of a base input, one or more source inputs, high-level causal variables targetted for intervention, and a counterfactual gold label that will be output by the network if the interchange intervention has the hypothesized effect on model behavior. This gold label is a counterfactual output of the high-level model we will align with the network. (See Appendix A.1 for details)
Optimize an orthogonal matrix to learn a distributed alignment for each high-level model that maximizes IIA using the training objective in Def. 3.4. We experiment with different hidden dimension sizes for our low-level model and different intervention site sizes (dimensionality of low-level subspaces) and locations (the layer where the intervention happens). (See Appendix A.2 for details)
Evaluate a baseline that brute-force searches through a discrete space of alignments and selects the alignment with the highest IIA. We search the space of alignments by aligning each high-level variable with groups of neurons in disjoint sliding windows. (See Appendix A.3 for details)
Evaluate the localist alignment “closest” to the learned distributed alignment. The rotation matrix for the localist alignment will be axis-aligned with the standard basis, possibly permuting and reflecting unit axes. (See Appendix A.4 for details)
Determine whether each distributed representation aligned with high-level variables can be decomposed into multiple representations that encode the identity of the input values to the variable’s causal mechanism. We do this by learning a second rotation matrix that decomposes learned distributed representation, holding the first rotation matrix fixed. (See Appendix A.5 for details)
The codebase used to run these experiments is athttps://github.com/atticusg/InterchangeInterventions/tree/zen. We have replicated the hierarchical equality experiment using the Pyvene library athttps://github.com/stanfordnlp/pyvene/blob/main/tutorials/advanced_tutorials/DAS_Main_Introduction.ipynb.
Hierarchical Equality Experiment
We now illustrate the power of DAS for analyzing networks designed to solve a hierarchical equality task. We concentrate on analyzing a trained feed-forward network.
A basic equality task is to determine whether a pair of objects are the same (). A hierarchical equality task is to determine whether a pair of pairs of objects have identical relations: . Specifically, the input to the task is two pairs of objects and the output is if both pairs are equal or both pairs are unequal and otherwise. For example, and are both assigned while is assigned .
We train a three-layer feed-forward network with ReLU activations to perform the hierarchical equality task. Each input object is represented by a randomly initialized vector. Specifically, our model has the following architecture where is the number of layers.
However, evaluating this high-level model alone is insufficient, as there are obviously many other high-level models of this task. To further contextualize our results, we also consider two alternatives: a high-level model where only the equality relation of the first pair is represented and a high-level model where the lone intermediate variable encodes the identity of the first input object (leaving all computation for the final step). These alternative high-level models also solve the task perfectly.
The IIA results achieved by the best alignment for each high-level model can be seen in Table 4.1. The best alignments found are with the ‘Both Equality Relations’ model that is widely assumed in the cognitive science literature. For all causal models, DAS learns a more faithful alignment (higher IIA) than a brute-force search through localist alignments. This result is most pronounced for ‘Both Equality Relations’, where DAS learns perfect or near-perfect alignments under a number of settings, whereas the best brute-force alignment achieves only 0.60 and the best localist alignment achieves only 0.73. Finally, the distributed representation of left equality could not be decomposed into a representation of the first argument identity. We see this in the very low performance of the ‘Identity Subspace of Left Equality’ results. This indicates that models are truly learning to encode an abstract equality relation, rather than merely storing the identities of the inputs.
In our second experiment, we analyze a BERT model fine-tuned on the Monotonicity Natural Language Inference (MoNLI) benchmark . A MoNLI example consists of a premise sentence and hypothesis sentence and the output label is entails when the premise makes the hypothesis true, and neutral otherwise. Two examples are in Figure 4. Every example is such that a single word in the premise sentence was changed to a hypernym (more general term) or hyponym (more specific term) to create the hypothesis. About half of MoNLI examples contain a negation that scopes over the word replacement site, and the remaining examples have no negation. When no negation is present, the label for a premise–hypothesis pair is the lexical relation. When negation is present, the label for a premise–hypothesis pair is the reverse of the lexical relation.
We fine-tune an uncased BERT-base model finetuned on the MultiNLI dataset .The parameters are provided by the Hugging Face library , downloaded from https://huggingface.co/ishan/bert-base-uncased-mnli. Our BERT model has 12 layers and 12 heads with a hidden dimension of 768. We concatenate the tokenized sequences of the premise sentence and hypothesis sentence with a token. Because of the size of the rotation matrix, we can’t look for distributed representations across all tokens; we look only at the representations of the token because the final classification is made from this token’s representation in the last layer.
2 High-Level Models
We use DAS to evaluate whether BERT fine-tuned on MoNLI will represent two boolean intermediate variables. The first is an indicator variable for negation, which is true if and only if negation is present in the premise and hypothesis. The second is a variable that is true if entails . This model is perhaps best expressed as a simple program (Figure 4). Again, we also consider two alternative high-level models to contextualize our results. One model represents only lexical entailment and not negation. The other represents the identity of the premise word .
& Hidden size Intervention size Layer 7 Layer 9 Layer 11 Layer 7 Layer 9 Layer 11 Layer 7 Layer 9 Layer 11 Layer 9 0.65 0.96 0.91 0.88 1.00 0.97 0.88 0.94 0.93 0.97 0.65 0.99 0.92 0.88 1.00 0.99 0.89 0.93 0.92 0.97 0.67 1.00 0.86 0.91 1.00 1.00 0.88 0.96 0.88 0.98 Brute-Force Search 0.60 0.56 0.52 0.64 0.64 0.57 0.50 0.51 0.54 - Localist Alignment 0.51 0.51 0.51 0.47 0.47 0.47 0.50 0.50 0.50 -
The IIA results achieved by the best alignment for each high-level model can be seen in Table 5.2. There is a perfect alignment between fine-tuned BERT and a symbolic algorithm with variables representing the presence of negation and the lexical entailment relation between and . In Table 5.2, this is shown by the perfect IIA for layer 9 and intervention size 256, meaning 256 non-standard basis dimensions of the token representation in layer 9 of BERT encode the relation between and and 256 other non-standard basis dimensions encode negation. Across all alignments and intervention types, DAS learns more faithful alignments (higher IIA) than a brute-force search through alignments, and no localist alignment comes close to the learned distributed alignments in terms of IIA.
However, the distributed representation of the lexical entailment relation between and can be nearly perfectly decomposed into two representations that encode the identity of the word and the identity of the word , respectively. This result is shown by the near perfect IIA in the final column of Table 5.2. This tells us that what appeared to be a representation of the lexical entailment, was in fact a “data structure” of two word identity representations.
We introduce distributed alignment search (DAS), a method to align interpretable causal variables with distributed neural representations. We learn distributed alignments that are more interpretable than localist alignments and do so with a gradient-descent based search method that improves upon the state-of-the-art brute-force search. In our two experiments, we discovered perfect alignments of distributed neural representations to binary high-level variables encoding simple equality and lexical entailment relations. However, when we investigated the substructure of these representations, we found that the lexical entailment representations could be decomposed into sub-representations of word identity. This highlights the need to investigate the causal substructure of neural representations. On the other hand, the presence of perfect representations of simple equality relations that cannot be decomposed into representations of the entities in the relations is a foundational result that should inform our understanding of how and when symbolic and connectionist architectures coexist.
This research is supported in part by grants from Open Philanthropy, Meta AI, Amazon, and the Stanford Institute for Human-Centered Artificial Intelligence (HAI).
Supplementary Materials
Appendix A Experimental Setup Details
For each task, we create training datasets for learning the rotation matrix of each high-level model. As defined in Definition 3.2, each input–output pair for training the rotation matrix consists of a base input that has two pairs of input values. Additionally, we have a set of source inputs mapping to interventions on different intermediate variables, and the corresponding counterfactual outputs (i.e., the updated outputs under interventions). Note that only for cases where there are multiple high-level intermediate variables involved, we sample more than one source input. For such cases, we randomly choose to interchange two variables together from two source inputs or swap a single variable from a single source input.
For our high-level models abstracting both equality relations and left equality relation, we sample a set of source inputs and interchange the equality relations of the corresponding shape pairs from the source inputs with the equality relations from the base input. For our high-level model abstracting the identity of the first shape, we sample a source input and interchange the first change from the source input with the base input.
Monotonicity NLI Experiments
For our high-level models abstracting negation or lexical entailment, we sample a set of source inputs and interchange the boolean value for negative or the value for lexical entailment from the source inputs with the base input. For our high-level model abstracting only the identity of replacing lexeme from the hypothesis sentence, we sample another hypothesis sentence from the one seen in training set and interchange its lexeme with the base input. To avoid cases where entailment labels are invalid (e.g., the entailment relation between “car” and “tree” is ambiguous), we specifically sample a valid English word that is either a hypernym or a hyponym of the lexeme item in the premise sentence, and from a new lexeme pair. Then, we construct a new pair of premise and hypothesis sentences by sampling a sentence template (i.e., a sentence with replaceable lexeme position such as “a man is talking to someone in a [lexeme]”) from the training dataset and replacing the lexeme items with new ones.
A.2 Reproducibility
We randomly generate 1.92M input–output pairs for training the model. We train our model for 10 epochs before reaching 100% training accuracy for the task. We also evaluate model performance on a hold-out testing set with unseen input-output pairs, and our model achieves 100% testing accuracy. For each high-level model, we then generate a training dataset for learning the rotation matrix. For each high-level model, we construct 640K such input–output pairs as our training data and 19.2K pairs as our testing data.
For both training phases, we use a batch size of 6.4K with a maximum training epoch of 10. We set the learning rate to 1e-3 with an early stop patient step set to 10K. Training with a single NVIDIA 2080 Ti RTX 11GB GPU takes less than ten minutes to converge. All datasets were balanced across the two labels during standard and interchange intervention training objectives. We run each experiment three times with distinct random seeds.
Monotonicity NLI Experiment
We randomly sample 10K examples from the original MoNLI dataset and use it to train our low-level models to solve MoNLI. We finetune our model for 5 epochs before reaching 100% training accuracy for the task. We also evaluate model performance on a hold-out testing set, and our model achieves 100% testing accuracy. For training and evaluating the rotation matrix of each high-level model, we create 24K examples as our training dataset for the first high-level model, and 10K for the rest two high-level models. For evaluation, we create 1.92K for the first high-level model, and 1K for the rest two high-level models.
We finetune our model for 5 epochs with a learning rate of before reaching 100% task accuracy with a batch size of 32. For the learning rotation matrix, we use a batch size of 64 with a learning rate of for a fixed epoch number of 5. Training with a single NVIDIA 2080 Ti RTX 11GB GPU takes less than ten minutes to converge for both training phases. We run each experiment three times with distinct random seeds.
A.3 Brute-Force Search Baseline
Without additional training, our brute-force search baseline finds the best IIA by searching over possible alignments as in Definition 3.5. For simple feed-forward networks, we map a high-level variable to a set of low-level variables within a sliding window with a size equal to the intervention size. We then incrementally search for the sliding window achieving the best IIA score starting from the first index of the intervened representation in the network. For Transformer-based networks, we avoid searching over all possible windows to make computation tractable, by only looking at windows with a starting index from of the token representation. Instead of targeting a specific set of layers in neural networks, we perform searches over all layers. Note that for the worst-case scenario, the number of hypotheses for the brute-force search approach becomes intractable and can be estimated as where is the total dimension size of the neural representation, and is the variable dimension size.
A.4 Localist Alignment Baseline
Without additional training, our localist alignment baseline finds a local optimal localist alignment matrix based on the learned rotation matrix. We pick the rotation matrix with the best IIA result from each category for evaluation. To find a localist alignment matrix, we follow Algorithm 1 to get our localist alignment matrix from any orthogonal matrix . We then use as our rotation matrix and evaluate IIA following our evaluation paradigm.
A.5 Subspace DAS
After learning a rotation matrix, we can fix it and learn another rotation matrix on top of it to do subspace high-level variable alignment. For instance, in the case of our MoNLI experiment, we fix the rotation matrix aligning the Lexical Entailment representation and further test whether we can learn another rotation matrix to align word identity. To achieve this, we initialize the first rotation matrix which aligns a larger subspace and freezes its weights along with the rest of the model. Then, we train another rotation matrix by taking the output representations from the first one with the same training objective as the first one as defined in Definition 3.4. The training data for the second rotation matrix is not the same as the first one, where we use the training data for the high-level model hypothesized to align with the subspace (e.g., the training data for the identity of first argument for the hierarchical equality task, and the training data for the identity of lexeme for the MoNLI task). Note that for both of our experiments, the subspace dimension is half of its parent subspace for simplicity.
Appendix B Runtime Comparison: Brute-force Search Baseline vs. DAS
Table 3 shows the runtime comparison between our method and brute-force search under the same settings for each task. Only our approach requires training. We underestimate the runtime for the brute-force search approach by only considering a limited set of possible alignments without exhaustively searching over the entire combination, which leads to intractable computations (See the BFS column of Table 3). The runtime of our approach can be further optimized if we deploy early stopping or optimized training data size, and it is invariant with the number of testing hypotheses.
Appendix C Remarks on Learned Rotation Matrix
Figure 5 shows the rotation in degree(s) of eigenvectorsThe eigenvectors of a rotation matrix are the vectors that remain unchanged after the rotation. of our learned rotation matrix for each task. We pick the best-performing oracle low-level model for each task for analyses. Our results suggest that learned rotations are not trivial, as the majority of basis vectors are rotated. These results suggest that the representations of high-level variables are highly distributed where direct probes over learned activation may fail to reveal the actual causal role of the representation effectively.
Appendix D Common Questions
In this section, we answer common questions that may be raised while reading this report.
Is the learned orthogonal matrix orthonormal?
Yes. We use the trainable orthogonal matrix implementation from PyTorch’s torch.nn.utils. parametrizations. It guarantees the resulting matrix is orthonormal when the rotation matrix is a full square matrix. Keeping the matrix orthonormal is crucial since it ensures we focus on rotation rather than scaling. Details can be found at https://pytorch.org/docs/stable/generated/torch.nn.utils.parametrizations.orthogonal.html.
How stable is the optimization process of the orthogonal matrix?
We rely on the default initialization of the orthogonal matrix in pytorch. The initialization step is important for finding the local optimal of the rotation matrix. In our experiment, we use random seeds and pick the best results out of our distinct runs to address this issue. However, we may consider different initialization schemes in the future.
Is an orthogonal matrix required to find distributed alignments?
In principle, the transformation is not required to be an orthogonal matrix. In fact, an orthogonal matrix assumes a linear transformation before aligning with a high-level variable, which may not be optimal if the aligning variable is represented in a non-linear sub-manifold of the representation space. In such cases, an orthogonal transformation results in imperfect interchange intervention accuracy, and an invertible and differentiable non-linear transformation may be more suitable (e.g., normalizing flow or invertible neural network). In practice, this transformation is computationally difficult to find, and the linear connections within neural networks also make them unlikely to be required to find alignments. We leave these investigations to future works.
What are the prerequisites to deploy this analysis method in practice?
We assume a partial or complete causal graph of the data generation process. Specifically, we assume to have interchangeable high-level variables defined for the causal graph. Additionally, we assume we can sample counterfactual data (i.e., base and source inputs where they differ in values of high-level variables) based on the causal graph.
How to interpret the result if the interchange intervention accuracy is not 100%?
When IIA is 100%, we rely on the graded notion of -on-average approximate causal abstraction , which directly coincides with IIA. More importantly, the relative IIA rankings between the high-level models also show which high-level model is a better approximation of the low-level model.
Does DAS scale with large foundation models?
Currently, the number of learnable parameters of the rotation matrix groups in polynomial time with the size of hidden representations. For instance, if our intervention site size is 512 in the lower-level model, the number of parameters of the rotation matrix is , which is about 0.26M. If we want to rotate concatenated token sequence embeddings of a BERT-BASE model in any layer, the number of parameters of the full rotation matrix is about 15.4B which becomes intractable for standard training infrastructure. To make computation tractable, DAS should be further reducible by representing only the aligned subspace, not the full rotation matrix. For instance, to find a 2-dim distributed representation within a 512-dimensional representation space, we approximately only need to learn parameters. In addition, we may use a low-rank approximation of the rotation matrix.
Practically, DAS transforms representations into an operatable state where interchange intervention results in interpretable model behaviors. DAS, itself, is a powerful tool for conducting causal abstraction analysis of a neural network.
Appendix E Task Performance & Interchange Intervention Accuracy Over Training Epochs
We additionally measure task performance (Task Acc.) as well IIA (Int. Acc.) of our alignments over training epochs for both seen training examples as well as unseen testing examples. Our results are shown from Figure 6 to Figure 11.