Transformer Interpretability Beyond Attention Visualization
Hila Chefer, Shir Gur, Lior Wolf
Introduction
Transformers and derived methods are currently the state-of-the-art methods in almost all NLP benchmarks. The power of these methods has led to their adoption in the field of language and vision . More recently, Transformers have become a leading tool in traditional computer vision tasks, such as object detection and image recognition . The importance of Transformer networks necessitates tools for the visualization of their decision process. Such a visualization can aid in debugging the models, help verify that the models are fair and unbiased, and enable downstream tasks.
The main building block of Transformer networks are self-attention layers , which assign a pairwise attention value between every two tokens. In NLP, a token is typically a word or a word part. In vision, each token can be associated with a patch . A common practice when trying to visualize Transformer models is, therefore, to consider these attentions as a relevancy score . This is usually done for a single attention layer. Another option is to combine multiple layers. Simply averaging the attentions obtained for each token, would lead to blurring of the signal and would not consider the different roles of the layers: deeper layers are more semantic, but each token accumulates additional context each time self-attention is applied. The rollout method is an alternative, which reassigns all attention scores by considering the pairwise attentions and assuming that attentions are combined linearly into subsequent contexts. The method seems to improve results over the utilization of a single attention layer. However, as we show, by relying on simplistic assumptions, irrelevant tokens often become highlighted.
In this work, we follow the line of work that assigns relevancy and propagates it, such that the sum of relevancy is maintained throughout the layers . While the application of such methods to Transformers has been attempted , this was done in a partial way that does not propagate attention throughout all layers.
Transformer networks heavily rely on skip connection and attention operators, both involving the mixing of two activation maps, and each leading to unique challenges. Moreover, Transformers apply non-linearities other than ReLU, which result in both positive and negative features. Because of the non-positive values, skip connections lead, if not carefully handled, to numerical instabilities. Methods such as LRP for example, tend to fail in such cases. Self-attention layers form a challenge since a naive propagation through these would not maintain the total amount of relevancy.
We handle these challenges by first introducing a relevancy propagation rule that is applicable to both positive and negative attributions. Second, we present a normalization term for non-parametric layers, such as “add” (\egskip-connection) and matrix multiplication. Third, we integrate the attention and the relevancy scores, and combine the integrated results for multiple attention blocks.
Many of the interpretability methods used in computer vision are not class-specific in practice, \ie, return the same visualization regardless of the class one tries to visualize, even for images that contain multiple objects. The class-specific signal, especially for methods that propagate all the way to the input, is often blurred by the salient regions of the image. Some methods avoid this by not propagating to the lower layers , while other methods contrast different classes to emphasize the differences . Our method provides the class-based separation by design and it is the only Transformer visualization method, as far as we can ascertain, that presents this property.
Explainability, interpretability, and relevance are not uniformly defined in the literature . For example, it is not clear if one would expect the resulting image to contain all of the pixels of the identified object, which would lead to better downstream tasks and for favorable human impressions, or to identify the sparse image locations that cause the predicted label to dominate. While some methods offer a clear theoretical framework , these rely on specific assumptions and often do not lead to better performance on real data. Our approach is a mechanistic one and avoids controversial issues. Our goal is to improve the performance on the acceptable benchmarks of the field. This goal is achieved on a diverse and complementary set of computer vision benchmarks, representing multiple approaches to explainability.
These benchmarks include image segmentation on a subset of the ImageNet dataset, as well as positive and negative perturbations on the ImageNet validation set. In NLP, we consider a public NLP explainability benchmark . In this benchmark, the task is to identify the excerpt that was marked by humans as leading to a decision.
Related Work
Many methods were suggested for generating a heatmap that indicates local relevancy, given an input image and a CNN. Most of these methods belong to one of two classes: gradient methods and attribution methods.
Gradient based methods are based on the gradients with respect to the input of each layer, as computed through backpropagation. The gradient is often multiplied by the input activations, which was first done in the Gradient*Input method . Integrated Gradients also compute the multiplication of the inputs with their derivatives. However, this computation is done on the average gradient and a linear interpolation of the input. SmoothGrad , visualizes the mean gradients of the input, and performs smoothing by adding to the input image a random Gaussian noise at each iteration. The FullGrad method offers a more complete modeling of the gradient by also considering the gradient with respect to the bias term, and not just with respect to the input. We observe that these methods are all class-agnostic: at least in practice, similar outputs are obtained, regardless of the class used to compute the gradient that is being propagated.
The GradCAM method is a class-specific approach, which combines both the input features and the gradients of a network’s layer. Being class-specific, and providing consistent results, this method is used by downstream applications, such as weakly-supervised semantic segmentation . However, the method’s computation is based only on the gradients of the deepest layers. The result, obtained by upsampling these low-spatial resolution layers, is coarse.
A second class of methods, the Attribution propagation methods, are justified theoretically by the Deep Taylor Decomposition (DTD) framework . Such methods decompose, in a recursive manner, the decision made by the network, into the contributions of the previous layers, all the way to the elements of the network’s input. The Layer-wise Relevance Propagation (LRP) method , propagates relevance from the predicated class, backward, to the input image based on the DTD principle. This assumes that the rectified linear unit (ReLU) non-linearity is used. Since Transformers typically rely on other types of applications, our method has to apply DTD differently. Other variants of attribution methods include RAP , AGF , DeepLIFT , and DeepSHAP . A disadvantage of some of these methods is the class-agnostic behavior observed in practice . Class-specific behavior is obtained by Contrastive-LRP (CLRP) and Softmax-Gradient-LRP (SGLRP) . In both cases, the LRP propagation results of the class to be visualized are contrasted with the results of all other classes, to emphasize the differences and produce a class-dependent heatmap. Our method is class-specific by construction and not by adding additional contrasting stages.
Methods that do not fall into these two main categories include saliency based methods , Activation Maximization and Excitation Backprop . Perturbation methods consider the change to the decision of the network, as small changes are applied to the input. Such methods are intuitive and applicable to black-box models (no need to inspect either the activations or the gradients). However, the process of generating the heatmap is computationally expensive. In the context of Transformers, it is not clear how to apply these correctly to discrete tokens, such as in text. Shapley-value methods have a solid theoretical justification. However, such methods suffer from a large computational complexity and their accuracy is often not as high as other methods. Several variants have been proposed, which improve both aspects .
Explainability for Transformers
There are not many contributions that explore the field of visualization for Transformers and, as mentioned, many contributions employ the attention scores themselves. This practice ignores most of the attention components, as well as the parts of the networks that perform other types of computation. A self-attention head involves the computation of queries, keys, and values. Reducing it only to the obtained attention scores (inner products of queries and keys) is myopic. Other layers are not even considered. Our method, in contrast, propagates through all layers from the decision back to the input.
LRP was applied for Transformers based on the premise that considering mean attention heads is not optimal due to different relevance of the attention heads in each layer . However, this was done in a limiting way, in which no relevance scores were propagated back to the input, thus providing partial information on the relevance of each head. We note that the relevancy scores were not directly evaluated, only used for visualization of the relative importance and for pruning less relevant attention heads.
The main challenge in assigning attributions based on attentions is that attentions are combining non-linearly from one layer to the next. The rollout method assumes that attentions are combined linearly and considers paths along the pairwise attention graph. We observe that this method often leads to an emphasis on irrelevant tokens since even average attention scores can be attenuated. The method also fails to distinguish between positive and negative contributions to the decision. Without such a distinction, one can mix between the two and obtain high relevancy scores, when the contributions should have cancelled out. Despite these shortcomings, the method was already applied by others to obtain integrated attention maps.
Abnar et al. present, in addition to rollout, a second method called attention flow. The latter considers the max-flow problem along the pair-wise attention graph. It is shown to be sometimes more correlated than the rollout method with relevance scores that are obtained by applying masking, or with gradients with respect to the input. This method is much slower and we did not evaluate it in our experiments for computational reasons.
We note this concurrent work did not perform an evaluation on benchmarks (for either rollout or attention-flow) in which relevancy is assigned in a way that is independent of the BERT network, for which the methods were employed. There was also no comparison to relevancy assignment methods, other than the raw attention scores.
Method
The method employs LRP-based relevance to compute scores for each attention head in each layer of a Transformer model . It then integrates these scores throughout the attention graph, by incorporating both relevancy and gradient information, in a way that iteratively removes the negative contributions. The result is a class-specific visualization for self-attention models.
Let be the number of classes in the classification head, and the class to be visualized. We propagate relevance and gradients with respect to class , which is not necessarily the predicted class. Following literature convention, we denote as the input of layer , where is the layer index in a network that consists of layers, is the input to the network, and is the output of the network.
Recalling the chain-rule, we propagate gradients with respect to the classifier’s output , at class , namely :
where the index corresponds to elements in , and corresponds to elements in .
We denote by the layer’s operation on two tensors and . Typically, the two tensors are the input feature map and weights for layer . Relevance propagation follows the generic Deep Taylor Decomposition :
where, similarly to Eq. 1, the index corresponds to elements in , and corresponds to elements in . Eq. 2 satisfies the conservation rule , \ie:
LRP assumes ReLU non-linearity activations, resulting in non-negative feature maps, where the relevance propagation rule can be defined as follows:
Non-linearities other that ReLU, such as GELU , output both positive and negative values. To address this, LRP propagation in Eq. 4 can be modified by constructing a subset of indices , resulting in the following relevance propagation:
In other words, we consider only the elements that have a positive weighed relevance.
2 Non parametric relevance propagation:
There are two operators in Transformer models that involve mixing of two feature map tensors (as opposed to a feature map with a learned tensor): skip connections and matrix multiplications (\egin attention modules). The two operators require the propagation of relevance through both input tensors. Note that the two tensors may be of different shapes in the case of matrix multiplication.
Given two tensors and , we compute the relevance propagation of these binary operators (\ie, operators that process two operands), as follows:
where and are the relevances for and respectively. These operations yield both positive and negative values.
The following lemma shows that for the case of addition, the conservation rule is preserved, \ie,
However, this is not the case for matrix multiplication.
Given two tensors and , consider the relevances that are computed according to Eq. 1. Then, (i) if layer adds the two tensors, \ie, then the conservation rule of Eq. 2 is maintained. (ii) if the layer performs matrix multiplication , then Eq. 2 does not hold in general.
(i) and (ii) are obtained from the output derivative of with respect to . In an add layer, and are independent of each other, while in matrix multiplication they are connected. A detailed proof of Lemma 3 is available in the supplementary. ∎
When propagating relevance of skip connections, we encounter numerical instabilities. This arises despite the fact that, by the conservation rule of the addition operator, the sum of relevance scores is constant. The underlying reason is that the relevance scores tend to obtain large absolute values, due to the way they are computed (Eq. 2). To see this, consider the following example:
where and are large positive numbers. It is easy to verify that . As can be seen, while the conservation rule is preserved, the relevance scores of and may explode. See supplementary for a step by step computation.
To address the lack of conservation in the attention mechanism due to matrix multiplication, and the numerical issues of the skip connections, our method applies a normalization to and :
Following the conservation rule (Eq. 3), and the initial relevance, we obtain for each layer .
The following lemma presents the properties of the normalized relevancy scores.
The normalization technique upholds the following properties: (i) it maintains the conservation rule, i.e.: , (ii) it bounds the relevance sum of each tensor such that:
3 Relevance and gradient diffusion
Let be a Transformer model consisting of blocks, where each block is composed of self-attention, skip connections, and additional linear and normalization layers in a certain assembly. The model takes as an input a sequence of tokens, each of dimension , with a special token for classification, commonly identified as the token [CLS]. outputs a classification probability vector of length , computed using the classification token. The self-attention module operates on a small sub-space of the embedding dimension , where is the number of “heads”, such that . The self-attention module is defined as follows:
Following the propagation procedure of relevance and gradients, each attention map has its gradients , and relevance , with respect to a target class , where is the layer that corresponds to the operation in Eq. 11 of block , and is the layer’s relevance.
For comparison, using the same notation, the rollout method is given by:
We can observe that the result of rollout is fixed given an input sample, regardless of the target class to be visualized. In addition, it does not consider any signal, except for the pairwise attention scores.
4 Obtaining the image relevance map
We consider only the tokens that correspond to the actual input, without special tokens, such as the [CLS] token and other separators. In vision models, such as ViT , the content tokens represent image patches. To obtain the final relevance map, we reshape the sequence to the patches grid size, \egfor a square image, the patch grid size is . This map is upsampled back to the size of the original image using bilinear interpolation.
Experiments
For the linguistic classification task, we experiment with the BERT-base model as our classifier, assuming a maximum of 512 tokens, and a classification token [CLS] that is used as the input to the classification head.
For the visual classification task, we experiment with the pretrained ViT-base model, which consists of a BERT-like model. The input is a sequence of all non-overlapping patches of size of the input image, followed by flattening and linear layers, to produce a sequence of vectors. Similar to BERT, a classification token [CLS] is appended at the beginning of the sequence and used for classification.
The baselines are divided into three classes: attention-maps, relevance, and gradient-based methods. Each has different properties and assumptions over the architecture and propagation of information in the network. To best reflect the performance of different baselines, we focus on methods that are both common in the explainability literature, and applicable to the extensive tests we report in this section, \egBlack-box methods, such as Perturbation and Shapely based methods, are computationally too expensive and inherently different from the proposed method. We briefly describe each baseline in the following section and the different experiments for each domain.
The attention-map baselines include rollout , following Eq. 16, which produces an explanation that takes into account all the attention-maps computed along the forward-pass. A more straightforward method is raw attention, \ieusing the attention map of block to extract the relevance scores. These methods are class-agnostic by definition.
Unlike attention-map based methods, the relevance propagation methods consider the information flow through the entire network, and not just the attention maps. These baselines include Eq. 4 and the partial application of LRP that follows . As we show in our experiments, the different variants of the LRP method are practically class-agnostic, meaning the visualization remains approximately the same for different target classes.
Evaluation settings For the visual domain, we follow the convention of reporting results for negative and positive perturbations, as well as showing results for segmentation, which can be seen as a general case of ”The Pointing-Game” . The dataset used is the validation set of ImageNet (ILSVRC) 2012, consisting of 50K images from 1000 classes, and an annotated subset of ImageNet called ImageNet-Segmentation , containing 4,276 images from 445 categories. For the linguistic domain, we follow ERASER and evaluate the reasoning for the Movies Reviews dataset, which consists of 1600/200/200 reviews for train/val/test. This task is a binary sentiment analysis task. Providing explanations for question answering and entailment tasks of the other datasets in ERASER, which require input sizes of more than 512 tokens (the limit of our BERT model), is left for future work.
The positive and negative perturbation tests follow a two-stage setting. First, a pre-trained network is used for extracting visualizations for the validation set of ImageNet. Second, we gradually mask out the pixels of the input image and measure the mean top-1 accuracy of the network. In positive perturbation, pixels are masked from the highest relevance to the lowest, while in the negative version, from lowest to highest. In positive perturbation, one expects to see a steep decrease in performance, which indicates that the masked pixels are important to the classification score. In negative perturbation, a good explanation would maintain the accuracy of the model, while removing pixels that are not related to the class. In both cases, we measure the area-under-the-curve (AUC), for erasing between of the pixels.
The two tests can be applied to the predicted or the ground-truth class. Class-specific methods are expected to gain performance in the latter case, while class-agnostic methods would present similar performance in both tests.
The segmentation tests consider each visualization as a soft-segmentation of the image, and compare it to the ground truth segmentation of the ImageNet-Segmentation dataset. Performance is measured by (i) pixel-accuracy, obtained after thresholding each visualization by the mean value, (ii) mean-intersection-over-union (mIoU), and (iii) mean-Average-Precision (mAP), which uses the soft-segmentation to obtain a score that is threshold-agnostic.
The NLP benchmark follows the evaluation setting of ERASER for rationales extraction, where the goal is to extract parts of the input that support the (ground truth) classification. The BERT model is first fine-tuned on the training set of the Movie Reviews Dataset and the various evaluation methods are applied to its results on the test set. We report the token-F1 score, which is best suited for per-token explanation (in contrast to explanations that extract an excerpt). To best illustrate the performance of each method, we consider a token to be part of the “rationale” if it is part of the top-k tokens, and show results for in steps of tokens. This way, we do not employ thresholding that may benefit some methods over others.
Qualitative evaluation Fig.3 presents a visual comparison between our method and the various baselines. As can be seen, the baseline methods produce inconsistent performance, while our method results in a much clearer and consistent visualization.
In order to show that our method is class-specific, we show in Fig. 3 images with two objects, each from a different class. As can be seen, all methods, except GradCAM, produce similar visualization for each class, while our method provides two different and accurate visualizations.
Perturbation tests Tab. 2 presents the AUC obtained for both negative and positive perturbation tests, for both the predicted and the target class. As can be seen, our method achieves better performance by a large margin in both tests. Notice that because rollout and raw attention produce constant visualization given an input image, we omit their scores in the target-class test.
Segmentation The segmentation metrics (pixel-accuracy, mAP, and mIoU) on ImageNet-segmentation are shown in Tab. 2. As can be seen, our method outperforms all baselines by a significant margin.
Language reasoning Fig. 3 depicts the performance on the Movie Reviews “rationales” experiment, evaluating for top-K tokens, ranging from to . As can be seen, while all methods benefit from increasing the amount tokens, our method consistently outperforms the baselines. See supplementary for a depiction of the obtained visualization.
Ablation study. We consider three variants of our method and present their performance on the segmentation and predicted class perturbation experiments. (i) Ours w/o , which modifies Eq. 13 s.t. we use instead of , (ii) , \iedisregarding rollout in Eq. 14, and using our method only on block , which is the block closest to the output, and (iii) which similar to (ii), only for block which is closer to the input.
As can be seen in Tab. 3 the ablation in which one removes the rollout component, \ie, Eq. 14, while keeping the relevance and gradient integration, and only considering the last attention layer, leads to a moderate drop in performance. Out of the two single block visualizations ((ii), and (iii)), the combined attention gradient and relevancy at the block, which is the closest to the output, is more informative than the block closest to the input. This is the same block that is being used for the raw-attention, partial LRP, and the GradCAM methods. The ablation that considers only this block outperforms these methods, indicating that the advantage of our method stems mostly from the combination of relevancy as we compute it and attention-map gradients.
Conclusions
The self-attention mechanism links each of the tokens to the [CLS] token. The strength of this attention link can be intuitively considered as an indicator of the contribution of each token to the classification. While this is intuitive, given the term “attention”, the attention values reflect only one aspect of the Transformer network or even of the self-attention head. As we demonstrate, both when using a fine-tuned BERT model for NLP and with the ViT model, attentions lead to fragmented and non-competitive explanations.
Despite this shortcoming and the importance of Transformer models, the literature with regards to interpretability of Transformers is sparse. In comparison to CNNs, multiple factors prevent methods developed for other forms of neural networks (not including the slower black-box methods) from being applied. These include the use of non-positive activation functions, the frequent use of skip connections, and the challenge of modeling the matrix multiplication that is used in self-attention.
Our method provides specific solutions to each of these challenges and obtains state-of-the-art results when compared to the methods of the Transformer literature, the LRP method, and the GradCam method, which can be applied directly to Transformers.
Acknowledgment
This project has received funding from the European Research Council (ERC) under the European Unions Horizon 2020 research and innovation programme (grant ERC CoG 725974). The contribution of the first author is part of a Master thesis research conducted at Tel Aviv University.
References
Appendix A Details of the various Baselines
As mentioned in Sec. 4, we consider the last attention layer (closest to the output) - namely . This results in a feature-map of size . Following the process described in Sec. 3.4, we take only the [CLS] token’s row (without the [CLS] token’s column), and reshape to the patches grid size . This results in a feature-map similar to the 2D feature-map used for GradCAM, where the number of channels, in this case, is , and the height and width are and . The reason we use the last attention layer is because of the sparse gradients issue described in Sec. 4.
raw-attention
The raw-attention method visualizes the last attention layer (closest to the output) - namely . It follows the process described in Sec. 3.4 to extract the final output.
LRP
In this method, we propagate relevance up to the input image, following the propagation rules of LRP (not our modified rules and normalizations).
partial-LRP
Following , we visualize an intermediate relevance map, more specifically, we visualize the last attention-map’s relevance, namely , using LRP propagation rules.
rollout
Appendix B Proofs for Lemmas
Given two tensors and , we compute the relevance propagation of binary operators (\ie, operators that process two operands) as follows:
where and are the relevances for and respectively.
The following lemma shows that for the case of addition, the conservation rule is preserved, i.e.,
However, this is not the case for matrix multiplication.
Given two tensors and , consider the relevances that are computed according to Eq. 1. Then, (i) if layer adds the two tensors, i.e., then the conservation rule of Eq. 2 is maintained. (ii) if the layer performs matrix multiplication , then Eq. 2 does not hold in general.
For part (i), we note that the number of elements in equals the number of elements in , therefore , and we can write Eq. 2 following the definition of :
note that, in this case, it is possible that .
As shown in the main text, while the sum of two tensors maintains the conservation rule, their values may explode. Consider , and , following the definition of we have:
To address the lack of conservation in the attention mechanism, which employs multiplication, and the numerical issues of the skip connections, our method applies a normalization to and :
The normalization technique upholds the following properties: (i) it maintains the conservation rule, i.e.: , (ii) it bounds the relevance sum of each tensor such that:
For part (ii) it is trivial to see that we weigh each tensor according to its relative absolute-value contribution:
Appendix C Visualizations - Multiple-class Images
Appendix D Visualizations - Single-class Images
Appendix E Visualizations - Text
In the following visualizations, we use the TAHV heatmap generator for text (https://github.com/jiesutd/Text-Attention-Heatmap-Visualization) to present the relevancy scores for each method, as well as the excerpts marked by humans. For methods that are class-dependent, we present the attributions obtained both for the ground truth class and the counter-factual class.
Evidently, our method is the only one that is able to present support for both sides, see panels (b,c) of each image. GradCAM often suffers from highlighting the evidence in the opposite direction (sign reversal), e.g., Fig. 13(g), in which the counter-factual explanation of GradCAM supports the negative, ground truth, sentiment and not the positive one.
Partial LRP (panels d,e) is not class-specific in practice. This provides it with an advantage in the quantitative experiments: Partial LRP highlights words with both positive and negative connotations from the same sentence, which better matches the behavior of the human annotators who are asked to mark complete sentences.
Notice that in most visualizations, it seems that the rollout method focuses mostly on the separation token [SEP], and fails to generate meaningful visualizations. This corresponds to the results presented in the quantitative experiments.
It seems from our results, e.g., Fig. 13(b,c) that the BERT tokenizer leads to unintuitive results. For example, “joyless” is broken down into “joy” and “less”, each supporting different sides of the decision.