This Looks Like That: Deep Learning for Interpretable Image Recognition

Chaofan Chen, Oscar Li, Chaofan Tao, Alina Jade Barnett, Jonathan Su, Cynthia Rudin

Introduction

How would you describe why the image in Figure 1 looks like a clay colored sparrow? Perhaps the bird’s head and wing bars look like those of a prototypical clay colored sparrow. When we describe how we classify images, we might focus on parts of the image and compare them with prototypical parts of images from a given class. This method of reasoning is commonly used in difficult identification tasks: e.g., radiologists compare suspected tumors in X-ray scans with prototypical tumor images for diagnosis of cancer . The question is whether we can ask a machine learning model to imitate this way of thinking, and to explain its reasoning process in a human-understandable way.

The goal of this work is to define a form of interpretability in image processing (this looks like that) that agrees with the way humans describe their own thinking in classification tasks. In this work, we introduce a network architecture – prototypical part network (ProtoPNet), that accommodates this definition of interpretability, where comparison of image parts to learned prototypes is integral to the way our network reasons about new examples. Given a new bird image as in Figure 1, our model is able to identify several parts of the image where it thinks that this part of the image looks like that prototypical part of some class, and makes its prediction based on a weighted combination of the similarity scores between parts of the image and the learned prototypes. In this way, our model is interpretable, in the sense that it has a transparent reasoning process when making predictions. Our experiments show that our ProtoPNet can achieve comparable accuracy with its analogous non-interpretable counterpart, and when several ProtoPNets are combined into a larger network, our model can achieve an accuracy that is on par with some of the best-performing deep models. Moreover, our ProtoPNet provides a level of interpretability that is absent in other interpretable deep models.

Our work relates to (but contrasts with) those that perform posthoc interpretability analysis for a trained convolutional neural network (CNN). In posthoc analysis, one interprets a trained CNN by fitting explanations to how it performs classification. Examples of posthoc analysis techniques include activation maximization , deconvolution , and saliency visualization . All of these posthoc visualization methods do not explain the reasoning process of how a network actually makes its decisions. In contrast, our network has a built-in case-based reasoning process, and the explanations generated by our network are actually used during classification and are not created posthoc.

Our work relates closely to works that build attention-based interpretability into CNNs. These models aim to expose the parts of an input the network focuses on when making decisions. Examples of attention models include class activation maps and various part-based models (e.g., ; see Table 1). However, attention-based models can only tell us which parts of the input they are looking at – they do not point us to prototypical cases to which the parts they focus on are similar. On the other hand, our ProtoPNet is not only able to expose the parts of the input it is looking at, but also point us to prototypical cases similar to those parts. Section 2.5 provides a comparison between attention-based models and our ProtoPNet.

Recently there have also been attempts to quantify the interpretability of visual representations in a CNN, by measuring the overlap between highly activated image regions and labeled visual concepts . However, to quantitatively measure the interpretability of a convolutional unit in a network requires fine-grained labeling for a significantly large dataset specific to the purpose of the network. The existing Broden dataset for scene/object classification networks is not well-suited to measure the unit interpretability of a network trained for fine-grained classification (which is our main application), because the concepts detected by that network may not be present in the Broden dataset. Hence, in our work, we do not focus on quantifying unit interpretability of our network, but instead look at the reasoning process of our network which is qualitatively similar to that of humans.

Our work uses generalized convolution by including a prototype layer that computes squared L2L^{2} distance instead of conventional inner product. In addition, we propose to constrain each convolutional filter to be identical to some latent training patch. This added constraint allows us to interpret the convolutional filters as visualizable prototypical image parts and also necessitates a novel training procedure.

Our work relates closely to other case-based classification techniques using k-nearest neighbors or prototypes , and very closely, to the Bayesian Case Model . It relates to traditional “bag-of-visual-words” models used in image recognition . These models (like our ProtoPNet) also learn a set of prototypical parts for comparison with an unseen image. However, the feature extraction in these models is performed by Scale Invariant Feature Transform (SIFT) , and the learning of prototypical patches (“visual words”) is done separately from the feature extraction (and the learning of the final classifier). In contrast, our ProtoPNet uses a specialized neural network architecture for feature extraction and prototype learning, and can be trained in an end-to-end fashion. Our work also relates to works (e.g., ) that identify a set of prototypes for pose alignment. However, their prototypes are templates for warping images and similarity with these prototypes does not provide an explanation for why an image is classified in a certain way. Our work relates most closely to Li et al. , who proposed a network architecture that builds case-based reasoning into a neural network. However, their model requires a decoder (for visualizing prototypes), which fails to produce realistic prototype images when trained on datasets of natural images. In contrast, our model does not require a decoder for prototype visualization. Every prototype is the latent representation of some training image patch, which naturally and faithfully becomes the prototype’s visualization. The removal of the decoder also facilitates the training of our network, leading to better explanations and better accuracy. Unlike the work of Li et al., whose prototypes represent entire images, our model’s prototypes can have much smaller spatial dimensions and represent prototypical parts of images. This allows for more fine-grained comparisons because different parts of an image can now be compared to different prototypes. Ming et al. recently took the concepts in and the preprint of an earlier version of this work, which both involve integrating prototype learning into CNNs for image recognition, and used these concepts to develop prototype learning in recurrent neural networks for modeling sequential data.

Case study 1: bird species identification

In this case study, we introduce the architecture and the training procedure of our ProtoPNet in the context of bird species identification, and provide a detailed walk-through of how our network classifies a new bird image and explains its prediction. We trained and evaluated our network on the CUB-200-2011 dataset of 200200 bird species. We performed offline data augmentation, and trained on images cropped using the bounding boxes provided with the dataset.

Figure 5 gives an overview of the architecture of our ProtoPNet. Our network consists of a regular convolutional neural network ff, whose parameters are collectively denoted by wconvw_{\text{conv}}, followed by a prototype layer gpg_{\mathbf{p}} and a fully connected layer hh with weight matrix whw_{h} and no bias. For the regular convolutional network ff, our model use the convolutional layers from models such as VGG-16, VGG-19 , ResNet-34, ResNet-152 , DenseNet-121, or DenseNet-161 (initialized with filters pretrained on ImageNet ), followed by two additional 1×11\times 1 convolutional layers in our experiments. We use ReLU as the activation function for all convolutional layers except the last for which we use the sigmoid activation function.

In our ProtoPNet, we allocate a pre-determined number of prototypes mkm_{k} for each class k∈{1,...,K}k\in\{1,...,K\} (1010 per class in our experiments), so that every class will be represented by some prototypes in the final model. Section S9.2 of the supplement discusses the choice of mkm_{k} and other hyperparameters in greater detail. Let Pk⊆P\mathbf{P}_{k}\subseteq\mathbf{P} be the subset of prototypes that are allocated to class kk: these prototypes should capture the most relevant parts for identifying images of class kk.

Finally, the mm similarity scores produced by the prototype layer gpg_{\mathbf{p}} are multiplied by the weight matrix whw_{h} in the fully connected layer hh to produce the output logits, which are normalized using softmax to yield the predicted probabilities for a given image belonging to various classes.

ProtoPNet’s inference computation mechanism can be viewed as a special case of a more general type of probabilistic inference under some reasonable assumptions. This interpretation is presented in detail in Section S2 of the supplementary material.

2 Training algorithm

The training of our ProtoPNet is divided into: (1) stochastic gradient descent (SGD) of layers before the last layer; (2) projection of prototypes; (3) convex optimization of last layer. It is possible to cycle through these three stages more than once. The entire training algorithm is summarized in an algorithm chart, which can be found in Section S9.3 of the supplement.

Stochastic gradient descent (SGD) of layers before last layer: In the first training stage, we aim to learn a meaningful latent space, where the most important patches for classifying images are clustered (in L2L^{2}-distance) around semantically similar prototypes of the images’ true classes, and the clusters that are centered at prototypes from different classes are well-separated. To achieve this goal, we jointly optimize the convolutional layers’ parameters wconvw_{\text{conv}} and the prototypes P={pj}j=1m\mathbf{P}=\left\{\mathbf{p}_{j}\right\}_{j=1}^{m} in the prototype layer gpg_{\mathbf{p}} using SGD, while keeping the last layer weight matrix whw_{h} fixed. Let D=[X,Y]={(xi,yi)}i=1nD=[\mathbf{X},\mathbf{Y}]=\{(\mathbf{x}_{i},y_{i})\}_{i=1}^{n} be the set of training images. The optimization problem we aim to solve here is:

The cross entropy loss (CrsEnt) penalizes misclassification on the training data. The minimization of the cluster cost (Clst) encourages each training image to have some latent patch that is close to at least one prototype of its own class, while the minimization of the separation cost (Sep) encourages every latent patch of a training image to stay away from the prototypes not of its own class. These terms shape the latent space into a semantically meaningful clustering structure, which facilitates the L2L^{2}-distance-based classification of our network.

In this training stage, we also fix the last layer hh, whose weight matrix is whw_{h}. Let wh(k,j)w_{h}^{(k,j)} be the (k,j)(k,j)-th entry in whw_{h} that corresponds to the weight connection between the output of the jj-th prototype unit gpjg_{\mathbf{p}_{j}} and the logit of class kk. Given a class kk, we set wh(k,j)=1w_{h}^{(k,j)}=1 for all jj with pj∈Pk\mathbf{p}_{j}\in\mathbf{P}_{k} and wh(k,j)=−0.5w_{h}^{(k,j)}=-0.5 for all jj with pj∉Pk\mathbf{p}_{j}\not\in\mathbf{P}_{k} (when we are in this stage for the first time). Intuitively, the positive connection between a class kk prototype and the class kk logit means that similarity to a class kk prototype should increase the predicted probability that the image belongs to class kk, and the negative connection between a non-class kk prototype and the class kk logit means that similarity to a non-class kk prototype should decrease class kk’s predicted probability. By fixing the last layer hh in this way, we can force the network to learn a meaningful latent space because if a latent patch of a class kk image is too close to a non-class kk prototype, it will decrease the predicted probability that the image belongs to class kk and increase the cross entropy loss in the training objective. Note that both the separation cost and the negative connection between a non-class kk prototype and the class kk logit encourage prototypes of class kk to represent semantic concepts that are characteristic of class kk but not of other classes: if a class kk prototype represents a semantic concept that is also present in a non-class kk image, this non-class kk image will highly activate that class kk prototype, and this will be penalized by increased (i.e., less negative) separation cost and increased cross entropy (as a result of the negative connection). The separation cost is new to this paper, and has not been explored by previous works of prototype learning (e.g., ).

Projection of prototypes: To be able to visualize the prototypes as training image patches, we project (“push”) each prototype pj\mathbf{p}_{j} onto the nearest latent training patch from the same class as that of pj\mathbf{p}_{j}. In this way, we can conceptually equate each prototype with a training image patch. (Section 2.3 discusses how we visualize the projected prototypes.) Mathematically, for prototype pj\mathbf{p}_{j} of class kk, i.e., pj∈Pk\mathbf{p}_{j}\in\mathbf{P}_{k}, we perform the following update:

The following theorem provides some theoretical understanding of how prototype projection affects classification accuracy. We use another notation for prototypes plk\mathbf{p}^{k}_{l}, where kk represents the class identity of the prototype and ll is the index of that prototype among all prototypes of that class.

Then after projection, the output logit for the correct class cc can decrease at most by Δmax⁡=m′log⁡((1+δ)(2−δ))\Delta_{\max}=m^{\prime}\log((1+\delta)(2-\delta)), and the output logit for every incorrect class k≠ck\neq c can increase at most by Δmax⁡\Delta_{\max}. If the output logits between the top-22 classes are at least 2Δmax⁡2\Delta_{\max} apart, then the projection of prototypes to their nearest latent training patches does not change the prediction of x\mathbf{x}.

Intuitively speaking, the theorem states that, if prototype projection does not move the prototypes by much (assured by the optimization of the cluster cost Clst), the prediction does not change for examples that the model predicted correctly with some confidence before the projection. The proof is in Section S1 of the supplement.

Note that prototype projection has the same time complexity as feedforward computation of a regular convolutional layer followed by global average pooling, a configuration common in standard CNNs (e.g., ResNet, DenseNet), because the former takes the minimum distance over all prototype-sized patches, and the latter takes the average of dot-products over all filter-sized patches. Hence, prototype projection does not introduce extra time complexity in training our network.

Convex optimization of last layer: In this training stage, we perform a convex optimization on the weight matrix whw_{h} of last layer hh. The goal of this stage is to adjust the last layer connection wh(k,j)w_{h}^{(k,j)} , so that for kk and jj with pj∉Pk\mathbf{p}_{j}\not\in\mathbf{P}_{k}, our final model has the sparsity property wh(k,j)≈0w_{h}^{(k,j)}\approx 0 (initially fixed at −0.5-0.5). This sparsity is desirable because it means that our model relies less on a negative reasoning process of the form “this bird is of class k′k^{\prime} because it is not of class kk (it contains a patch that is not prototypical of class kk).” The optimization problem we solve here is: min⁡wh1n∑i=1nCrsEnt(h∘gp∘f(xi),yi)+λ∑k=1K∑j:pj∉Pk∣wh(k,j)∣\min_{w_{h}}\frac{1}{n}\sum_{i=1}^{n}\textrm{CrsEnt}(h\circ g_{\mathbf{p}}\circ f(\mathbf{x_{i}}),\mathbf{y_{i}})+\lambda\sum_{k=1}^{K}\sum_{j:\mathbf{p}_{j}\not\in\mathbf{P}_{k}}|w_{h}^{(k,j)}|. This optimization is convex because we fix all the parameters from the convolutional and prototype layers. This stage further improves accuracy without changing the learned latent space or prototypes.

3 Prototype visualization

Given a prototype pj\mathbf{p}_{j} and the training image x\mathbf{x} whose latent patch is used as pj\mathbf{p}_{j} during prototype projection, how do we decide which patch of x\mathbf{x} (in the pixel space) corresponds to pj\mathbf{p}_{j}? In our work, we use the image patch of x\mathbf{x} that is highly activated by pj\mathbf{p}_{j} as the visualization of pj\mathbf{p}_{j}. The reason is that the patch of x\mathbf{x} that corresponds to pj\mathbf{p}_{j} should be the one that pj\mathbf{p}_{j} activates most strongly on, and we can find the patch of x\mathbf{x} on which pj\mathbf{p}_{j} has the strongest activation by forwarding x\mathbf{x} through a trained ProtoPNet and upsampling the activation map produced by the prototype unit gpjg_{\mathbf{p}_{j}} (before max-pooling) to the size of the image x\mathbf{x} – the most activated patch of x\mathbf{x} is indicated by the high activation region in the (upsampled) activation map. We then visualize pj\mathbf{p}_{j} with the smallest rectangular patch of x\mathbf{x} that encloses pixels whose corresponding activation value in the upsampled activation map from gpjg_{\mathbf{p}_{j}} is at least as large as the 9595th-percentile of all activation values in that same map. Section S7 of the supplement describes prototype visualization in greater detail.

4 Reasoning process of our network

Figure 5 shows the reasoning process of our ProtoPNet in reaching a classification decision on a test image of a red-bellied woodpecker at the top of the figure. Given this test image x\mathbf{x}, our model compares its latent features f(x)f(\mathbf{x}) against the learned prototypes. In particular, for each class kk, our network tries to find evidence for xx to be of class kk by comparing its latent patch representations with every learned prototype pj\mathbf{p}_{j} of class kk. For example, in Figure 5 (left), our network tries to find evidence for the red-bellied woodpecker class by comparing the image’s latent patches with each prototype (visualized in “Prototype” column) of that class. This comparison produces a map of similarity scores towards each prototype, which was upsampled and superimposed on the original image to see which part of the given image is activated by each prototype. As shown in the “Activation map” column in Figure 5 (left), the first prototype of the red-bellied woodpecker class activates most strongly on the head of the testing bird, and the second prototype on the wing: the most activated image patch of the given image for each prototype is marked by a bounding box in the “Original image” column – this is the image patch that the network considers to look like the corresponding prototype. In this case, our network finds a high similarity between the head of the given bird and the prototypical head of a red-bellied woodpecker (with a similarity score of 6.4996.499), as well as between the wing and the prototypical wing (with a similarity score of 4.3924.392). These similarity scores are weighted and summed together to give a final score for the bird belonging to this class. The reasoning process is similar for all other classes (Figure 5 (right)). The network finally correctly classifies the bird as a red-bellied woodpecker. Section S3 of the supplement provides more examples of how our ProtoPNet classifies previously unseen images of birds.

5 Comparison with baseline models and attention-based interpretable deep models

The accuracy of our ProtoPNet (with various base CNN architectures) on cropped bird images is compared to that of the corresponding baseline model in the top of Table 1: the first number in each cell gives the mean accuracy, and the second number gives the standard deviation, over three runs. To ensure fairness of comparison, the baseline models (without the prototype layer) were trained on the same augmented dataset of cropped bird images as the corresponding ProtoPNet. As we can see, the test accuracy of our ProtoPNet is comparable with that of the corresponding baseline (non-interpretable) model: the loss of accuracy is at most 3.5%3.5\% when we switch from the non-interpretable baseline model to our interpretable ProtoPNet. We can further improve the accuracy of ProtoPNet by adding the logits of several ProtoPNet models together. Since each ProtoPNet can be understood as a “scoring sheet” (as in Figure 5) for each class, adding the logits of several ProtoPNet models is equivalent to creating a combined scoring sheet where (weighted) similarity with prototypes from all these models is taken into account to compute the total points for each class – the combined model will have the same interpretable form when we combine several ProtoPNet models in this way, though there will be more prototypes for each class. The test accuracy on cropped bird images of combined ProtoPNets can reach 84.8%84.8\%, which is on par with some of the best-performing deep models that were also trained on cropped images (see bottom of Table 1). We also trained a VGG19-, DenseNet121-, and DenseNet161-based ProtoPNet on full images: the test accuracy of the combined network can go above 80%80\% – at 80.8%80.8\%, even though the test accuracy of each individual network is 72.7%72.7\%, 74.4%74.4\%, and 75.7%75.7\%, respectively. Section S3.1 of the supplement illustrates how combining several ProtoPNet models can improve accuracy while preserving interpretability.

Moreover, our ProtoPNet provides a level of interpretability that is absent in other interpretable deep models. In terms of the type of explanations offered, Figure 5 provides a visual comparison of different types of model interpretability. At the coarsest level, there are models that offer object-level attention (e.g., class activation maps ) as explanation: this type of explanation (usually) highlights the entire object as the “reason” behind a classification decision, as shown in Figure 5(a). At a finer level, there are numerous models that offer part-level attention: this type of explanation highlights the important parts that lead to a classification decision, as shown in Figure 5(b). Almost all attention-based interpretable deep models offer this type of explanation (see the bottom of Table 1). In contrast, our model not only offers part-level attention, but also provides similar prototypical cases, and uses similarity to prototypical cases of a particular class as justification for classification (see Figure 5(c)). This type of interpretability is absent in other interpretable deep models. In terms of how attention is generated, some attention models generate attention with auxiliary part-localization models trained with part annotations (e.g., ); other attention models generate attention with “black-box” methods – e.g., RA-CNN uses another neural network (attention proposal network) to decide where to look next; multi-attention CNN uses aggregated convolutional feature maps as “part attentions.” There is no explanation for why the attention proposal network decides to look at some region over others, or why certain parts are highlighted in those convolutional feature maps. In contrast, our ProtoPNet generates attention based on similarity with learned prototypes: it requires no part annotations for training, and explains its attention naturally – it is looking at this region of input because this region is similar to that prototypical example. Although other attention models focus on similar regions (e.g., head, wing, etc.) as our ProtoPNet, they cannot be made into a case-based reasoning model like ours: the only way to find prototypes on other attention models is to analyze posthoc what activates a convolutional filter of the model most strongly and think of that as a prototype – however, since such prototypes do not participate in the actual model computation, any explanations produced this way are not always faithful to the classification decisions. The bottom of Table 1 compares the accuracy of our model with that of some state-of-the-art models on this dataset: “full” means that the model was trained and tested on full images, “bb” means that the model was trained and tested on images cropped using bounding boxes (or the model used bounding boxes in other ways), and “anno.” means that the model was trained with keypoint annotations of bird parts. Even though there is some accuracy gap between our (combined) ProtoPNet model and the best of the state-of-the-art, this gap may be reduced through more extensive training effort, and the added interpretability in our model already makes it possible to bring richer explanations and better transparency to deep neural networks.

6 Analysis of latent space and prototype pruning

In this section, we analyze the structure of the latent space learned by our ProtoPNet. Figure 5(a) shows the three nearest prototypes to a test image of a Florida jay and of a cardinal. As we can see, the nearest prototypes for each of the two test images come from the same class as that of the image, and the test image’s patch most activated by each prototype also corresponds to the same semantic concept as the prototype: in the case of the Florida jay, the most activated patch by each of the three nearest prototypes (all wing prototypes) indeed localizes the wing; in the case of the cardinal, the most activated patch by each of the three nearest prototypes (all head prototypes) indeed localizes the head. Figure 5(b) shows the nearest (i.e., most activated) image patches in the entire training/test set to three prototypes. As we can see, the nearest image patches to the first prototype in the figure are all heads of black-footed albatrosses, and the nearest image patches to the second prototype are all yellow stripes on the wings of golden-winged warblers. The nearest patches to the third prototype are feet of some gull. It is generally true that the nearest patches of a prototype all bear the same semantic concept, and they mostly come from those images in the same class as the prototype. Those prototypes whose nearest training patches have mixed class identities usually correspond to background patches, and they can be automatically pruned from our model. Section S8 of the supplement discusses pruning in greater detail.

Case study 2: car model identification

In this case study, we apply our method to car model identification. We trained our ProtoPNet on the Stanford Cars dataset of 196196 car models, using similar architectures and training algorithm as we did on the CUB-200-2011 dataset. The accuracy of our ProtoPNet and the corresponding baseline model on this dataset is reported in Section S6 of the supplement. We briefly state our performance here: the test accuracy of our ProtoPNet is comparable with that of the corresponding baseline model (≤3%\leq 3\% difference), and that of a combined network of a VGG19-, ResNet34-, and DenseNet121-based ProtoPNet can reach 91.4%91.4\%, which is on par with some state-of-the-art models on this dataset, such as B-CNN (91.3%91.3\%), RA-CNN (92.5%92.5\%), and MA-CNN (92.8%92.8\%).

Conclusion

In this work, we have defined a form of interpretability in image processing (this looks like that) that agrees with the way humans describe their own reasoning in classification. We have presented ProtoPNet – a network architecture that accommodates this form of interpretability, described our specialized training algorithm, and applied our technique to bird species and car model identification.

Supplementary Material and Code: The supplementary material and code are available at https://github.com/cfchen-duke/ProtoPNet.

This work was sponsored in part by a grant from MIT Lincoln Laboratory to C. Rudin.

References