Differentiable Patch Selection for Image Recognition
Jean-Baptiste Cordonnier, Aravindh Mahendran, Alexey Dosovitskiy, Dirk Weissenborn, Jakob Uszkoreit, Thomas Unterthiner
Introduction
High-resolution imagery has become ubiquitous nowadays: both consumer devices and specialized sensors routinely capture images and videos with resolution in tens of megapixels. Processing these high-quality images with computer vision models remains challenging: analyzing the images at full resolution can be prohibitively computationally expensive, while simply downsampling them before processing may remove important fine details and substantially hurt performance. It would be desirable to save compute, while retaining the capability to recognize fine details.
Compute can be saved by exploiting the following property of many practical vision tasks: not all parts of the image are equally important for finding the answer. Figure 1 shows examples of tasks where only a small fraction of the full image needs to be processed in detail. Being able to quickly discard uninformative parts of the image would have several benefits. It would reduce the overall computational and memory complexity of the model, and the regions of interest could be processed in more detail and by a more powerful model than otherwise.
Determining which parts of the image to retain and which to discard is usually nontrivial and highly task dependent. In some applications the solution might be as simple as taking the center crop of the image, but in most cases relevant regions need to be detected first. For instance, in a self-driving car setting, it would be permissible to ignore the sky, but all traffic signs in sight should be correctly identified and must not be ignored. One may formulate this as follows: Given a regular grid of equally sized image patches, decide for each patch whether to process or discard it. This decision is however discrete, which makes it unsuitable for end-to-end learning.
To overcome this limitation, inspired by the work of Katharopoulos & Fleuret , we formulate patch selection as a ranking problem, where per-patch relevance scores are predicted by a small ConvNet and the Top scoring patches are selected for downstream processing. We make this end-to-end trainable with backpropagation using the perturbed maximum method of Berthet et. al . We present this as a generic module for patch selection. Our approach is most effective when the majority of patches in the image are irrelevant to the target, but the model a priory does not know where in the image the important patches are present. Hence, we do not aim to achieve image coverage such as in semantic segmentation and object detection.
In the remainder of this paper, we will formulate patch selection for image recognition as a Top-K selection problem in section 3, apply the perturbed maximum method to construct an end-to-end model trainable via backpropagation in section 4.2, and demonstrate wide applicability of this method via empirical results in three different domains: (1) street sign recognition, (2) inter-patch relationship reasoning on synthetic data, and (3) fine-grained classification without using object/part bounding box annotations during training and evaluation (Section 5).
Related Work
: Several computer vision methods extract regions of interest from the image. Two stage object detection approaches, for instance, select regions of interest using region proposal networks or hand crafted heuristics . Selected regions are later processed by a separate stage of the model. These methods use the non-differentiable RoI-Pooling or the differentiable RoI-Align . Such architectures require bounding box supervision to train large scale object detection models, whereas our experiments focus on a simpler setting and aim to train with weak supervision using only a single class label per image.
Soft attention
: In order to attend to specific parts of an image, an alternative approach is to occlude parts of the input by generating attention masks . While this helps the model focus on relevant features , become more interpretable , or include external data such as image captions , the models will typically still process the whole image on a fixed input resolution. Thus they do not lead to any efficiency gains. Another approach would be to process several image resolutions in parallel and use an attention mechanism to pick features from them. It is also possible to employ adhoc losses to extract meaningful patches .
Multiple-Instance Learning
: A number of works use attention to solve Multiple-Instance Learning (MIL) problems, which are especially common in medical imaging, where images tend to be very large . Here the goal is to label a set of related input samples, such as slices of organ scans or large images that are decomposed into patches. While, for example, the method of Ilse et. al can be used to identify the most relevant patches, this is not leveraged to make computation more efficient, as all image patches are processed in equal detail by this method.
Sequential “glimpses”
: There is a long line of work that sequentially processes a sequence of patches (“glimpses”), from a network, until they settle on the most relevant ones . These methods often rely on Reinforcement Learning to train non-differentiable attention mechanisms , which typically makes them difficult to train. The Spatial Transformers on the other hand can be deployed as a differentiable attention mechanism, for example for fine-grained recognition of bird species. These have been applied sequentially to extract several regions of interest as part of recurrent neural networks. Training spatial transformers on large images can, however, be difficult because the gradients with respect to the transformation parameters are an accumulation of gradients of sub-pixel bilinear interpolation which can be very local. Angles et. al overcome these limitations to train a multiple instance spatial transformers by lifting non-differentiable Top-K by introducing an auxiliary function that creates a heat-map given a set of interest points. Our method, on the other hand, avoids these limitations by computing gradients with respect to all patches in every backward step.
Attention Sampling
Differentiable Top-K
: The subset sampling operation can be implemented by expanding the Gumbel-Softmax trick or based on optimal-transport formulations for ranking and sorting , the latter of which was recently made significantly faster by Blondel et. al . Our work uses perturbed optimizers to make a Top-K differentiable, which we found performed better than the Sinkhorn operator .
Model overview
Our model processes high resolution images by scoring, selecting then processing, some regions of interest. As illustrated in Figure 2, the model consists of a scorer network , a patch selection module , a feature network and an aggregation network , where , and are learnable parameters. Attention sampling (ATS) is a special case of this where aggregation is an average pool and patch selection is implemented using discrete sampling. To describe the model in one sentence: The scorer scores patches, the patch selection module selects a subset of those, the feature network computes embeddings for each patch, and the aggregation network combines these embeddings to make a final prediction. These modules are described in more detail below.
Several prior works can be cast in this manner: Spatial transformers have localisation networks that are similar to our scorer network, albeit instead of scoring patches they predict localization parameters such as an affine transformation. Thus scoring and patch selection are effectively done together. Their grid sampler can be interpreted as a patch extraction module. Attention sampling sample patches (with or without replacement) according to normalized scores output by a scorer network and use a ResNet for feature extraction . The aggregation network is restricted to taking the average of the patch embeddings to be able to back propagate through the discrete sampling operation. Vision transformer : the patch selection module is exhaustive and extracts all patches in the image with a stride . The feature network and the aggregation network are merged into a large transformer that process all the flattened patches with self-attention and output the representation of the CLS token or an average of the tokens’ representations.
Patch selection as differentiable Top-K
2 Differentiable Top-K
To learn the parameters of the scorer network using backpropagation, we need to differentiate through the patch selection method. We employ the perturbed maximum method for this purpose. Given a non-differentiable module, whose forward pass can be represented as a linear program of the form
with inputs , optimization variable , and convex polytope constraint set , the perturbed maximum method defines a differentiable module with forward and backward operations as below.
Sample uniform Gaussian noise and perturb the input , to generate several perturbed inputs. Solve the linear program for each of these, and then average their results.
where is a hyper-parameter. In practice one computes the expectation using an empirical mean with independent samples for . We fix in all our experiments and tune . This does not require one to solve linear programs in every forward pass, instead as the linear program is chosen to be equivalent to Top-K, we run the Top-K algorithm times, one for each perturbed input, which is very fast in practice.
Backward:
Following , the Jacobian associated with the above forward pass is
The equations above have been simplified for the special case of normal distributed Other distributions could be used. Please see for details.
Patch selection as Top-K with sorted indices is equivalent to the following linear program.
Note how has the same shape as the concatenation of indicator vectors in the previous subsection. They serve the same purpose. The first conditions encourages that is an assignment where each of the columns has a total weight of one. The last condition results in sorting the indices. This linear program has infinitely many optimal solutions, one of which is the required integer solution corresponding to the index-sorted “Top-K” operator.
Note that in theory, noise should be applied to . The equivalence between this linear program and Top-K crucially relies on all columns of being identical. We therefore apply noise directly to . Our experiments show that this departure from theory works in practice. We normalize the scorer output to lie in $\epsilon=10^{-5}$ to avoid any division by zero.
While Top-K for each perturbed input will result in one-hot indicators , their perturbed average may be far from one-hot. Thus early in training, when the scores are still non-decisive, extracted patches resemble a weighted average over all image patches (Figure 3). Mixing images like this might contribute to model generalization . Another side benefit is that backpropagated gradients take into account all image patches and can update their weights from the very first step. A discrete sampler on the other hand is solving a combinatorial search over patches and gradients may be non-informative until the right patches have been sampled consistently for a few iterations of training. The average can, however, become meaningless if the Top-K solver returned a different permutation of the most salient patches for each perturbed input. Sorting of indices overcomes this issue and is an integral part of our “Top-K” operator.
For efficiency reasons we use hard top-K during inference. Hard top-K processes 3-10% more images per second because only a single Top-K operation has to be performed (instead of perturbed repetitions) and patch extraction by slicing the tensor is more efficient than weighted combination with the indicator vectors. Furthermore, hard Top-K is deterministic which might be a desirable property at inference. However, using hard top-K at inference results in a train-test gap. To bridge this we linearly decay to zero during training. At , no noise is added and the differentiable top-k operation is numerically identical to hard top-K. The gradients flowing into the scorer network vanish at .
We also experimented with another differentiable Top-K formulation based on Sinkhorn operators . We found the above perturbed-optimizer formulation to give superior results. Sinkhorn based experiments are therefore deferred to the appendix (Section D).
Experiments and results
Differentiable patch selection is a generic tool that can benefit a variety of computer vision problems. In this work, we focus on selecting a small number of patches from high resolution images for image classification.
Our first task is to recognize speed limits signs in large images. This is a key task to enable autonomous driving and is a natural fit for our method since the relevant pixels in the image are very localized and the model must rely on the high resolution images to read speed indications at significant distances. We use the Swedish traffic signs dataset , replicating the setup of for their Attention Sampling (ATS) method. ATS uses a subset of the dataset consisting of 747 training images and 684 test images of dimension pixels, and the goal is to classify whether each image contains a limit sign of 50, 70 or 80 kilometers per hour or no speed limit. We apply the same data augmentation as , specifically a random translation and a random affine color scaling per image. For a fair comparison, we use ATS’s scorer: the scorer is a 4-layers CNN followed by a stride 8 max-pooling layer. We apply the scorer on a downscaled version of the image. We obtain scores for candidate patches of dimensions and select of them. The feature network is the modified thin ResNet used by ATS. We use mean-pooling to aggregate per-patch representations.
We ablate against not having a patch selection procedure and directly apply the feature network CNN to the full image. Beside wasting computation over constant parts of images (e.g. the blue sky), this small ResNet completely overfits the training set: reaching 100% training accuracy while not achieving better accuracy than just predicting the majority class on test. It seems that the inductive bias introduced by patch extraction is crucial in this very low data regime.
We manually investigated the mistakes made by our model. Anecdotally the scoring network was able to extract the most relevant patches in most cases. Its only failure mode was a tendency to extract false-positive patches that exhibit the same colors as the traffic signs in question. Most miss-classifications are likely due to either the feature- or the aggregation-network. One could potentially further improve our results by pre-training the feature network.
2 Inter-patch reasoning
The work most closely related to ours, ATS, requires averaging representations obtained from sampled patches. We hypothesize that such mean aggregation limits expressivity required for modelling relationships between extracted patches. We investigate this using a synthetic dataset inspired by MegaMNIST , but which goes beyond scattered MNIST numbers on a Megapixel image to instead consist of billiard balls as presented in Figure 6. Each image contains four to eight randomly colored balls randomly placed on the table. Ball numbers are sampled uniformly from and face the camera. We ensure that balls do not completely obstruction each other. We define the following classification task: report the higher of two numbers extracted from the leftmost and rightmost balls. All the other balls can be ignored. This task exhibits three interesting properties: (i) the information on the image is very localized (around the balls), (ii) downsampling the image severely degrades the readability of ball numbers, (iii) the task may be difficult to solve using a simple mean-pooling approach.
Using the Kubrichttps://github.com/google-research/kubric software, we generated 20k images of size 10001000. We split them into 8k samples for training, 2k for validation, and 10k for test. The generated dataset is available for downloadhttp://storage.googleapis.com/gresearch/ptokp_patch_selection/billiard.tar.xz. Code for our experiments is available on GitHubhttps://github.com/google-research/google-research/tree/master/ptopk_patch_selection.
We apply our model to this task. Architecture details are as follows. The scorer is a 4-layer CNN and processes the downscaled image of size where the balls are visible but the numbers cannot be read. The feature network is ResNet18 . We compare different aggregation schemes: mean-pooling, max-pooling, and a small Transformer . The latter consists of 3 self-attention layers with 8 heads, taking the sequence of patch representations as input augmented with a learned additive positional encoding. We compare against ATS as well as a simple ResNet18 baseline. All the methods are trained using Adam with decoupled weight decay of and tuned learning rate . We repeated each experiment several times, but noted that not all runs were successful. For ATS, only 1 out of 4 runs were able to meaningfully solve the problem by reaching an accuracy of 53.25 %, while the remaining three attempts performed no better than predicting the majority class. We also tried concatenating $$ normalized fixed positional encodings as two extra channels in the input (‘concat. position’) as the task requires the model to consider ball position. This affords a fairer comparison between our Transformer variant and the ATS and ResNet baselines. Positional encoding considerably improved baseline performances but were not as effective as the transformer. With concatenated position encodings, 3 out of 4 ATS runs meaningfully solved the problem while 1 run always predicted the majority class. This failure mode of majority class prediction also happened to 2 out of 9 runs of our model when using max-pooling aggregation. Both mean-pooling and transformer aggregation were relatively stable. We report robust estimates of performance (median and median absolute deviation) in Table 2.
As the results show, a standard CNN is able to solve this task, but the training time increase by 24 % compared to differentiable Top-K with a transformer on top. This difference in running time would be even more pronounced in higher resolutions (see supplementary material). One may reduce CNN training time by downsampling the input image. As demonstrated in the results, this does not work, as the numbers become too small to be readable. ATS is not able to solve this task reliably, most likely because of mean pooling per-patch embeddings. Our own method also performs poorly when using mean-pooling, but when using transformers and max-pooling aggregation we are able to solve this task with high accuracy, outperforming all other methods. We compare the three aggregation schemes using a box-whiskers plot in fig. 7.
3 Fine-grained bird classification
Another way to extract salient regions of the image is to use a teacher-student approach to learn how to rank patches (NTS-Net). We argue that NTS-Net’s ranking loss is essential only due to the non-differentiability of their patch selection method, and show here how their model can be simplified using our differentiable Top-K module. We explore this using the Caltech-UCSD Birds (CUB-200) dataset , a common dataset for fine-grained image classification. This dataset is a bird classification task with 11,788 images from 200 bird species.
We briefly describe the training setup of NTS-Net: The original image is resized and cropped to , from which ResNet activations are computed. A scorer network processes these activations to score patches at multiple scales and aspect ratios. The scorer is trained using a ranking loss which we skip here because it is not required in our formulation. top scoring (Top-K) patches are selected post greedy non-maxima suppression, resized to , and encoded using the same ResNet backbone. Call these activations . NTS-Net tries to predict the class label from , each of , and a concatenation of all these embeddings . The model minimize a softmax cross entropy loss on each prediction head.
In this section, the emphasis is not on efficiently processing high resolution input but rather on being able to process salient parts of the object at a level of detail that is computationally infeasible when processing the entire input. We adapt out architecture to closely match NTS-Net, modulo four differences. First, NTS-Net uses greedy non-maximum suppression (NMS) in its patch selection module. Hard greedy NMS cannot be made differentiable using the perturbed maximum method as it does not correspond to the optimum of any linear program. Non-greedy-NMS, on the other hand, can be easily incorporated into the constraint set (eq. 5). It would be computationally infeasible to solve the resulting linear program times in each forward pass. One could learn a network to do NMS to overcome this problem. We instead resort to making our scorer network more expressive using Squeeze-Excitation layers , so that it may learn to select non-overlapping patches if necessary. This way of doing feature modulation enables global communication between all spatial locations, allowing us to model a crude form of global context. Second, we use entropy regularization with a coefficient of as in other experiments above Regularization is applied on the softmax of concatenated unnormalized scores. The final RPN outputs are normalized to lie between $R_{I}+\frac{\left(R_{P_{1}}+\cdots+R_{P_{K}}\right)}{K}2048$ dimensional.
We tuned optimization hyper-parameters and regularization strength by splitting the training set into 5000 training images and 994 validation images. This was to avoid meta-overfitting on test data. We took the best hyper-parameter combination from this train-val split and then trained the model on the entire 5994 training images and tested on the official test set. Test results are reported in Table 3.
We note that our method performs slightly worse than the NTS-Net baseline. This may be due to the following: (a) Our model may be more prone to overfitting. It can select patches that directly optimize the training loss. We combat this using ‘mean’ aggregation which improves performance as shown in the last two rows in Table 3. (b) We report average performance across 5 seeds. Our best performing ‘mean’-aggregation model achieves 87.3%. (c) Despite using SE modules in the scorer network, our model still selects overlapping patches around the bird head. This likely also makes overfitting worse. Please see supplementary material for qualitative results.
Discussion and future work
We introduced a differentiable Top-K module for patch selection in large images. This enabled end-to-end learning of the patch selection module in a task dependent manner. We observed competitive performance on a task where information relevant for classification was very localized within the image. We significantly advanced the state-of-the-art in a problem requiring inter-patch relationship reasoning, as evidence by our use of a Transformer module on patch embeddings in the billiard balls dataset. Our method allows for arbitrary downstream neural processing of patches without using auxiliary losses for training the scorer network.
We are delighted by the above progress but note that there is still significant room for improvement and identify the following areas for future work. The patch selection problem is a chicken-and-egg problem. One can pick optimal patches by first knowing the contents of the image, but to efficiently know the contents of the image one should pick the right patches. This makes it inherently difficult to properly account for global context. Training patch selection models remains difficult. Some failure modes include predicting the majority class, or the scorer becoming very confident of the wrong patches either early in training or diverging to this behavior in the middle of training. More adaptive optimization might help improve learning stability. Patch selection lead to less overfitting on the Traffic street signs dataset. However, the same model overfits on the CUB-200 dataset. The scorer was able to identify patches associated with watermarks and background in order to memorize the training set. Fortunately when combined with the NTS-Net framework, overfitting was considerably reduced. Future work could address this by training on larger dataset or using self-supervised learning.
Acknowledgements
We thank Thomas Kipf, Francesco Locatello, Georg Heigold, Klaus Greff, Sindy Löwe, and Quentin Berthet for helpful discussions.
References
Appendix A Supplemental Material
The supplementary material consists of the following: performance trade-offs associated with patch sampling versus running a CNN on the entire high resolution image (Appendix B), some theoretical discussions regarding the LP formulation of index-sorted Top-K in (Appendix C), an experiment quantifying the effect of decaying to 0 during training in (Appendix D), results with another differentiable top-K method we experimented with before developing our approach (Appendix E), hyper-parameter details for all our experiments (Appendix F), qualitative results (Appendix G), and a PyTorch implementation of the perturbed Top-K module (Appendix H).
Appendix B Speed Improvements by Sampling Patches
We study the speed improvement that can be gained at inference by using our patch extraction model compared to running a model on the full image. We compare the number of samples processed per second at inference on a single V100 GPU in Figure 8. If the useful information for recognition is localized within than 10% of the pixels, which corresponds to extracting patches of size on a MegaPixel image, processing only the relevant regions with the same network (ResNet50) allows a 5 fold speed up. In this case, roughly 28 % of the inference time is spent on the feature network, while most of the remaining time is spent calculating the scores for the individual patches. The Top-K operation’s influence on runtime is minimal. This should be put in contrast with the common alternative of running a significantly smaller ResNet18 on the full size image which is not faster on images than our approach running a ResNet50 on the extracted patches. Using hard Top-K at inference instead of the differentiable version is consistently faster on a V100 GPU.
Appendix C Linear Program
To understand why the LP, in equations 4 and 5, corresponds to index sorted top-K, let us focus on the integral solutions . In this scenario, each column of is a one-hot indicator vector selecting exactly one of the scores in . Furthermore, the objective function is the sum of selected scores. Thus the optimal solution is to select the highest scores.
which is a contradiction to the original assumption.
Appendix D Decaying σ𝜎\sigma to 00
We linearly decay perturbation magnitude to in all our models. We ablate this design choice on the billiard balls dataset. These experiments use the transformer aggregation head. All other hyper-parameters are kept the same. In fig. 9 and table 4 we see that using hard top-k at inference leads to higher accuracy (). Decaying perturbation magnitude to zero further improves performance ().
Appendix E Differentiable Sinkhorn
Another approach to make Top-K differentiable was proposed by Xie et. al and relies on the optimal transport formulation of Top-K proposed by Cuturi et. al . We implemented the forward and backward pass following Algorithm 3 of . We report the results for traffic sign recognition in Table 5 with similar setting as in Section 5.1. This approach gives good results but suffers when using discrete Top-K at inference.
Appendix F Experimentation Details
For all experiments except the Fine-grained bird classification ones, our scorer network consisted of a CNN with 4 convolutional layers of kernel size with stride 1 and “valid” padding. The number of feature maps was 8, 16, 32 and 1, respectively. Every convolution except for the last was followed by a Relu activation function. The last convolution was followed by an max pooling of stride 8.
Our feature network was different depending on the dataset: On the Swedish Traffic Signs dataset, we used the same feature network as ATS , which is a narrow ResNet18 with 16 filters instead of the usual 64, 128, 256, 512. On the billiard balls dataset we used a standard ResNet18. On the CUB-200 dataset we use the same feature extractor as NTS-Net, that is a ResNet50.
We used the ADAM-W optimizer for Traffic signs and billiard balls dataset. We used SGD with momentum , similar to NTS-Net, for CUB-200. The exact details of our optimizers can be seen in Table 6. We found that weight decay coupled with momentum updates was better to reduce overfitting on CUB-200.
In our adaptation of the NTS-Net, we added Squeeze Excitation layers inside the region proposal network (RPN). This was meant to compensate for the lack of non-maximum suppression in our patch selection module. The RPN for both our method and the baseline NTS-Net are shown in figure 10. Lastly, with regards to data pre-processing, instead of normalizing pixels by a pre-computed pixel mean and pixel standard deviation, we re-scaled pixel values to lie between $$ during training. This was to match the value range we used for pre-training our ResNet50 backbone on ImageNet ILSVRC12. Other than this minor change, we used the same data augmentation as NTS-Net.
Appendix G Qualitative Results
We visualize patches extracted by the model on the CUB-200 dataset, on the test split, in fig. 11. This particular model is the best performing seed among the 5 random repeats we averaged to report the number for in the main manuscript. It uses ‘mean’ aggregation and achieves top-1 accuracy. The images visualized are not cherry picked. We see that the model is able to focus on the bird and captures the region around the eyes and torso. The lack of NMS is evident as the model often looks at patches with high overlap.