Cross Attention Network for Few-shot Classification

Ruibing Hou, Hong Chang, Bingpeng Ma, Shiguang Shan, Xilin Chen

Introduction

Few-shot classification aims at classifying unlabeled samples (query set) into unseen classes given very few labeled samples (support set). Compared to traditional classification, few-shot classification has two main challenges: One is unseen classes, i.e., the non-overlap between training and test classes; The other is the low-data problem, i.e., very few labeled samples for the test unseen classes.

Solving few-shot classification problem requires the model trained with seen classes to generalize well to unseen classes with only few labeled samples. A straightforward approach is fine-tuning a pre-trained model using the few labeled samples from the unseen classes. However, it may cause severe overfitting. Regularization and data augmentation can alleviate but cannot fully solve the overfitting problem. Recently, meta-learning paradigm is widely used for few-shot learning. In meta-learning, the transferable meta-knowledge, which can be an optimization strategy , a good initial condition , or a metric space , is extracted from a set of training tasks and generalizes to new test tasks. The tasks in the training phase usually mimic the settings in the test phase to reduce the gap between training and test settings and enhance the generalization ability of the model.

While promising, few of them pay enough attention to the discriminability of the extracted features. They generally extract features from the support classes and unlabeled query samples independently, as a result, the features are not discriminative enough. For one thing, the test images in the support/query set are from unseen classes, thus their features can hardly attend to the target objects. To be specific, for a test image containing multiple objects, the extracted feature may attend to the objects from seen classes which have large number of labeled samples in the training set, while ignore the target object from unseen class. As illustrated in Fig. 1 (c) and (d), for the two images from the test class curtain, the extracted features only capture the information of the objects that are related to the training classes, such as person or chair in Fig. 1 (a) and (b). For another, the low-data problem makes the feature of each test class not representative for the true class distribution, as it is obtained from very few labeled support samples. In a word, the independent feature representations may fail in few-shot classification.

In this work, we propose a novel Cross Attention Network (CAN) to enhance the feature discriminability for few-shot classification. Firstly, Cross Attention Module (CAM) is introduced to deal with the unseen class problem. The cross attention idea is inspired by the human few-shot learning behavior. To recognize a sample from unseen class given a few labeled samples, human tends to firstly locate the most relevant regions in the pair of labeled and unlabeled samples. Similarly, given a class feature map and a query sample feature map, CAM generates a cross attention map for each feature to highlight the target object. Correlation estimation and meta fusion are adopted to achieve this purpose. In this way, the target object in the test samples can get attention and the features weighted by the cross attention maps are more discriminative. As shown in Fig. 1 (e), the extracted features with CAM can roughly localize the regions of target object curtain. Secondly, we introduce a transductive inference algorithm that utilizes the entire unlabeled query set to alleviate the low-data problem. The proposed algorithm iteratively predicts the labels for the query samples, and selects pseudo-labeled query samples to augment the support set. With more support samples per class, the obtained class features can be more representative, thus alleviating the low-data problem.

Experiments are conducted on multiple benchmark datasets to compare the proposed CAN with existing few-shot meta-learning approaches. Our method achieves new state-of-the art results on all dataset, which demonstrates the effectiveness of CAN.

Related Work

Few-Shot Classification. On the basis of the availability of the entire unlabeled query set, few-shot classification can be divided into two categories: inductive and transductive few-shot classification. In this work, we mainly explore the few-shot approaches based on meta-learning.

Inductive Few-shot Learning has been a well studied area in recent years. One promising way is the meta-learning paradigm. It usually trains a meta-learner from a set of tasks, which extracts meta-knowledge to transfer into new tasks with scarce data. Meta learning approaches for few-shot classification can be roughly categorized into three groups. Optimization-based methods designed the meta-learner as an optimizer that learned to update model parameters . Further, the works learned a good initialization so that the learner could rapidly adapt to novel tasks within a few optimization steps. Parameter-generating based methods usually designed the meta-learner as a parameter predicting network. Metric-learning based methods learned a common feature space where categories can distinguish with each other based on a distance metric. For example, Matching Network produced a weighted nearest neighbor classifier. Prototypical Network performed nearest neighbor classification with learned class features (prototypes). The works improved the prototypical network with a learnable similarity metric , a task adaptive metric , or a image-to-class local metric .

Our proposed framework belongs to metric-learning based method. Different from existing metric-learning based methods which extracted the support and query sample features independently, our method exploits the semantic relevance between support and query features to highlight the target object. Although the parameter-generating based methods also consider the relationship between support and query samples, these approaches require an additional complex parameter prediction network. With less overload, our approach outperforms these methods by a large margin.

Transductive Algorithm. Transductive few-shot classification is firstly introduced in , which constructed a graph on the support set and the entire query set, and propagated labels within the graph. However, the method required a specific architecture, making it less universal. Inspired by the self-training strategy in semi-supervised learning , we propose a simper and more general transductive few-shot algorithm, which explicitly augments the labeled support set with unlabeled query samples to achieve more representative class features. Moreover, the proposed transductive algorithm can be directly applied to the existing models, e.g., prototypical network , matching network , and relation network .

Attention Model. Attention mechanisms aim to highlight important local regions to extract more discriminative features. It has achieved great success in computer vision applications, such as image classification , image caption and visual question answering . In image classification, SENet proposed a channel attention block to boost the representational power of a network. Woo et al. further integrated the channel and spatial attention modules to a block. However, these blocks are not effective for few-shot classification. We argue that they localized the important regions of the test images only based on the priors of the training classes, which could not generalize to the test images from unseen classes. For example, as shown in Fig. 1, since curtain is not in the training classes, above blocks would attend to the foregrounds of seen classes, such as the person or the chair, rather than the curtain. In image caption, the attention blocks usually used the last generated words to search for related regions in the image to generate the next word. And in visual question answering, the attention blocks used the questions to localize the related regions in the image to answer. For few-shot image classification, in this paper, we design a meta-learner to compute the cross attention between support (or class) and query feature maps, which helps to locate the important regions of the target object and enhance the feature discriminability.

Cross Attention Module

Problem Define. Few-shot classification usually involves a training set, a support set and a query set. The training set contains a large number of classes and labeled samples. The support set of few labeled samples and the query set of unlabeled samples share the same label space, which is disjoint with that of the training set. Few-shot classification aims to classify the unlabeled query samples given the training set and support set. If the support set consists of CC classes and KK labeled samples per class, the target few-shot problem is called CC-way KK-shot.

Following , we adopt the episode training mechanism, which has been demonstrated as an effective approach for few-shot learning. The episodes used in training simulate the settings in test. Each episode is formed by randomly sampling CC classes and KK labeled samples per class as the support set S={(xas,yas)}a=1ns\mathcal{S}=\{\left(x^{s}_{a},y^{s}_{a}\right)\}_{a=1}^{n_{s}} (ns=C×Kn_{s}=C\times K), and a fraction of the rest samples from the CC classes as the query set Q={(xbq,ybq)}b=1nq\mathcal{Q}=\{\left(x^{q}_{b},y^{q}_{b}\right)\}_{b=1}^{n_{q}}. And we denote Sk\mathcal{S}^{k} as the support subset of the kthk^{th} class. How to represent each support class Sk\mathcal{S}^{k} and query sample xbqx^{q}_{b} and measure the similarity between them is a key issue for few-shot classification.

CAM Overview. In this work, we resort to metric-learning to obtain proper feature representations for each pair of support class and query sample. Different from existing methods which extract the class and query features independently, we propose Cross Attention Module (CAM) which can model the semantic relevance between the class feature and query feature, thus draw attention to the target objects and benefit the subsequent matching.

Complexity Analysis. The time and space cost of CAM is mainly on correlation layer. The time complexity of CAM is O(h2w2c)O(h^{2}w^{2}c) and the space complexity is O(hwc)O(hwc), which both varies with the size of input feature map. So we insert CAM after the last convolutional layer to avoid excessive cost.

Cross Attention Network

Model Training via Optimization. CAN is trained via minimizing the classification loss on the query samples of training set. The classification module consists of a nearest neighbor and a global classifier. The nearest neighbor classifier classifies the query samples into CC support classes based on pre-defined similarity measure. To obtain precise attention maps, we constrain each position in the query feature maps to be correctly classified. Specifically, for each local query feature qibq^{b}_{i} at ithi^{th} position, the nearest neighbor classifier produces a softmax-like label distribution over CC support classes. The probability of predicting qibq^{b}_{i} as kthk^{th} class is:

where (Qˉkb)i(\bar{Q}_{k}^{b})_{i} denotes the feature vector in the ithi^{th} spatial position of Qˉkb\bar{Q}_{k}^{b}, and GAP is the global average pooling operation to get the mean class feature. Note that Qˉkb\bar{Q}_{k}^{b} and Qˉjb\bar{Q}_{j}^{b} represent the query sample xbqx^{q}_{b} from somewhat different views as they are correlated with different support classes. In Eq. 4, the cosine distance dd is calculated in the feature space generated by CAM. The nearest neighbor classification loss is then defined as the negative log-probability according to the true class label ybq∈{1,2,…,C}y^{q}_{b}\in\{1,2,\dots,C\}:

Inductive Inference. In inductive inference phase, the embedding module is directly used for a novel task to extract the class and query feature maps. Then each pair of class and query feature maps are fed into CAM to get the attention weighted features. The global averaging pooling is then performed to the outputs of CAM to get the mean class and query features. Finally, the label y^bq\hat{y}^{q}_{b} for a query sample xbqx^{q}_{b} is predicted by finding the nearest mean class feature under cosine distance metric:

Transductive Inference. In few-shot classification task, each class has very few labeled samples, so the class feature can hardly represent the true class distribution. In order to alleviate the problem, we propose a simple and effective transductive inference algorithm which utilizes the unlabeled query samples to enrich the class feature.

Specifically, we firstly utilize the initial class feature map PkP^{k} to predict the labels {y^bq}b=1nq\{\hat{y}^{q}_{b}\}_{b=1}^{n_{q}} of the unlabeled query samples {xbq}b=1nq\{x^{q}_{b}\}_{b=1}^{n_{q}} using Eq. 7. Then, we define a label confidence criterion using the cosine distance between the query sample xbqx^{q}_{b} and its nearest class neighbor: cbq=min⁡kd(GAP(Qˉkb),GAP(Pˉbk))c^{q}_{b}=\min_{k}d(\textit{GAP}(\bar{Q}_{k}^{b}),\textit{GAP}(\bar{P}_{b}^{k})). The lower the value cbqc^{q}_{b}, the higher the confidence of the predicted label {y^bq}\{\hat{y}^{q}_{b}\}. Based on this criterion, we can obtain a candidate set D={(xbq,y^bq)∣sb=1,xbq∈Q}\mathcal{D}=\{(x^{q}_{b},\hat{y}^{q}_{b})|s_{b}=1,x^{q}_{b}\in\mathcal{Q}\}, where sb∈{0,1}s_{b}\in\{0,1\} denotes the selection indicator for the query sample xbqx^{q}_{b}. The selection indicator s∈{0,1}nqs\in\{0,1\}^{n_{q}} is determined by the top tt confident query samples: s=arg⁡min⁡∣∣s∣∣0=t∑b=1nqsbcbqs=\arg\min_{||s||_{0}=t}\sum_{b=1}^{n_{q}}s_{b}c^{q}_{b}. Finally, the candidate set D\mathcal{D} along with the support set S\mathcal{S} is used to generate a more representative class feature map (Pk)∗(P^{k})^{*}:

Here Dk={(xbq,y^bq)∣xbq∈D,y^bq=k}\mathcal{D}^{k}=\{(x^{q}_{b},\hat{y}^{q}_{b})|x^{q}_{b}\in\mathcal{D},\hat{y}^{q}_{b}=k\}. (Pk)∗(P^{k})^{*} is then used to re-estimate the pseudo label for each query sample. We repeat above process for a certain number of iterations. And the number of selected samples in the candidates set D\mathcal{D} is gradually increased with a fixed ratio in each iteration. In this way, we can progressively enrich the class features to be more representative and robust.

Experiments

Datasets. We use miniImageNet which is a subset of ILSVRC-12 . It contains 100100 classes with 600600 images per class. We use the standard split following : 6464 classes for training, 1616 for validation and 2020 for testing. We also use tieredImageNet dataset , a much larger subset of ILSVRC-12 . It contains 3434 categories and 608608 classes in total. These are divided into 2020 categories (351351 classes) for training, 66 categories (9797 classes) for validation, and 88 categories (160160 classes) for testing, as in .

Experimental setting. We experiment our approach on 55-way 11-shot and 55-way 55-shot settings. For a CC-way KK-shot setting, the episode is formed with CC classes and each class includes KK support samples, and 66 and 1515 query samples are used for training and inference respectively. When inference, 20002000 episodes are randomly sampled from the test set. We report the average accuracy and the corresponding 95%95\% confidence interval over the 20002000 episodes.

Implementation details. Pytorch is used to implement all our experiments on one NVIDIA 1080Ti GPU. Following , we use ResNet-12 network as our embedding module. The input images size is 84×8484\times 84. During training, we adopt horizontal flip, random crop and random erasing as data augmentation. SGD is used as the optimizer. Each mini-batch contains 88 episodes. The model is trained for 8080 epochs, with each epoch consisting of 1,2001,200 episodes. For miniImageNet, the initial learning rate is 0.10.1 and decreased to 0.0060.006 and 0.00120.0012 at 6060 and 7070 epochs, respectively. For tieredImageNet, the initial learning rate is set to 0.10.1 with a decay factor 0.10.1 at every 2020 epochs. The temperature hyperparameter (τ\tau in Eq. 3) is set to 0.0250.025, the reduction ratio in the meta-learner is set to 66, and the weight hyperparameter (λ\lambda) in the overall loss function is set to 0.50.5. For the transductive algorithm, the selected number of query samples in the first iteration (tt) is set to 3535, and the number of iterations and enlarging factor of candidate set are both set to 22. All hyperparameters are cross-validated in the validation sets and fixed afterwards in all experiments.

2 Comparison with State-of-the-arts

Tab. 1 compares our method with existing few-shot methods on miniImageNet and tieredImageNet. The comparative methods are categorized into four groups, i.e., optimization-based methods (O), parameter-generating methods (P), metric-learning methods (M), and transductive methods (T). Our method outperforms the optimization-based methods . It is noted that the optimization-based methods need fine-tuning on the target task, making the classification time consuming. On the contrary, our method requires no model updating solves the tasks in an feed-forward manner, which is much faster and simpler than above methods and has better results.

Our method performs better than the parameter-generating methods , with an improvement up to 7%7\%. These approaches generate the parameters of the feature extractor based on the support set and extract the query features adaptively. However, these methods suffer from the high dimensionality of the parameter space. Instead, our method uses a cross attention module to adaptively extract the support and query features, which is computationally lightweight and achieves a better performance. Our method belongs to the metric-learning methods. Existing metric-based methods extract features of support and query samples independently, making the features attend to non-target objects. Instead, our CAN highlights the target object regions and gets more discriminative features. Compared to TADAM , CAN with almost the same number of parameters achieves 5%5\% higher performance on 1-shot, which demonstrates the superiority of our cross attention module.

In the transductive setting, CAN with transduction (CAN+T) outperforms the prior work TPN by a large margin, up to 8%8\% and 5%5\% improvements on 1-shot and 5-shot respectively. TPN uses a graph network to propagate the labels of the support set to the query set. In contrast, our algorithm selects the top confident query samples to augment the support set, which can explicitly alleviate the low-data problem. In addition, our transductive algorithm can be easily applied to other few-shot learning models, e.g., matching network , prototypical network and relation network .

Time complexity comparison. Tab. 1 further compares the time cost of our method to others. Some methods use a 4-layer ConvNet as the backbone thus take relatively lower time cost. Even though, our CAN is still comparable even superior to these methods in term of time cost, with a performance improvement up to 10%10\%. The others use the same backbone as CAN, but require following up modules such as model update per task , gradient-based parameter generation , or expensive condition generation , which all incur more time overhead than CAM. Overall, Tab. 1 shows that CAN outperforms other methods without excessive overhead.

3 Ablation Study

In this subsection, we empirically show the effectiveness of each component of CAN. We firstly introduce two baselines to be used for comparison. In R12-proto , the features from embedding module are directly fed to the nearest neighbor classifier and the model is trained with nearest neighbor classification loss. In R12-proto-ac, the only difference from R12-proto is that R12-proto-ac has an additional logit head for global classification (the normal 6464-way classification in miniImageNet case) and the model is trained with the joint of global and nearest neighbor classification loss.

Influence of global classification. The comparison results are shown in Tab. 2. By comparing R12-proto-ac to R12-proto, we can find large improvements on both 1-shot (5.8%5.8\%) and 5-shot (7.7%). We further try another meta-learner matching network (MN) We re-implement matching network with ResNet-12 as backbone on miniImageNet., and the proposed joint learning schema improves MN from 55.29%55.29\% to 59.14%59.14\% on 1-shot setting and 67.74%67.74\% to 73.81%73.81\% on 5-shot setting. The consistent improvements demonstrate the effectiveness of the joint leaning schema. We argue that the global classification loss provides regularization on the embedding module and forces it to perform well on two decoupled tasks, nearest neighbor classification and global classification.

Influence of cross attention module. By comparing our CAN to R12-pro-ac, we observe consistent improvements on both 1-shot and 5-shot scenarios. The reason is that when using the cross attention module, our model is able to highlight the relevant regions and extract more discriminative feature. The performance gap also provides evidence that (1) conventionally independently extracted features tend to focus on non-target region and produce inaccurate similarities. (2) cross attention module can help to highlight target regions and reduce such inaccuracy with small overhead.

Influence of meta-learner in CAM. To verify the effectiveness of the meta-learner in CAM, we develop two variants of CAM without meta-learner. Specifically, one variant named CAN-NoML-1 sets the kernel ww (shown in Fig. 2 (b)) to be a fixed mean kernel, i.e., performing global average pooling on the correlation map RR to get the attention maps AA. The other variant, CAM-NoML-2, sets the kernel ww to a vanilla learnable convolutional kernel that remains the same for all input samples. As shown in Tab. 2, both variants outperform R12-proto-ac consistently, which further demonstrates the effectiveness of the proposed cross attention mechanism. The improvements of CAN-NoML-1 shows the mean of correlators can roughly estimate the relevant semantic information, which furthers verifies the reasonability of our designed meta-leaner. As seen, CAN outperforms both variants. The improvement can be attributed to the meta-learning schema which learns to adaptively generate the kernel ww according to the input pair of feature maps.

Influence of transductive inference algorithm. As shown in Tab. 2, CAN+T greatly improves CAN especially in 1-shot where the low-data problem is more serious. To further verity its effectiveness, we apply it to other few-shot models, i.e., matching network , prototypical network and relation network . We re-implement these models using the code provided by to ensure a fair comparison. As shown in Tab. 3, our algorithm consistently improves the performance of these models, which demonstrates its generalization ability. Nevertheless, the improvements to these models are inferior to CAN. We argue that CAN can predict more precise pseudo labels for query samples and augment the support set more effectively, thus leading to better performance.

Complexity comparisons. To illustrate the cost of CAN, we report the number of parameters (PN), the number of floating-point operations (GFLOPs) and the average CPU inference time (CIT) for a 5-way 1-shot and 5-way 5-shot task with 15 query samples per class. As shown in Tab. 2, CAN introduces negligible parameters (the parameters W1W_{1} and W2W_{2} of the meta-learner in CAM) and small computational overhead. For example, CAN requires 101.81101.81 GFLOPs for 5-way 1-shot, corresponding to only 0.25%0.25\% relative increase over original R12-proto-ac. Notably, the correlation map in CAM can be worked out by one matrix multiplication, which occupies less time in GPU libraries. The transductive inference algorithm also introduces small computational overhead (0.37%0.37\% on 1-shot and 0.31%0.31\% on 5-shot) since it directly utilizes the extracted embedding features to regenerate the class feature and only passes the lightweight CAM again.

4 Visualization Analysis

To qualitatively evaluate the proposed cross attention mechanism, we compare the class activation maps visualization results of CAN to other meta-learners, RN , MAML and TADAM . As shown in Fig. 4 (a), the features of RN usually contain non-target objects since it lacks an explicit mechanism for feature adaptation. MAML performs gradient-based adaptation, which makes the model merely learn some conspicuous discriminative features in the support images without deeping into the intrinsic characteristic of the target objects. As shown in Fig. 4 (b), MAML attends to ship for the groenendael support image to better distinguish it from the golden retriever category, resulting in a confusing location and misclassification of the groenendael category. TADAM performs task-dependent adaptation and applies the same adaptive parameters to all query images of a task, thus it is difficult to locate different target objects for different categories. As shown in Fig. 4 (c), TADAM mistakenly attends to the dog for worm fence query image. In contrast, CAN processes the query samples with different adaptive parameters, which allows it to focus on the different target objects for different categories shown in Fig. 4 (d).

Conclusion

In this paper, we proposed a cross attention network for few-shot classification. Firstly, a cross attention module is designed to model the semantic relevance between class and query features. It can adaptively localize the relevant regions and generate more discriminative features. Secondly, we propose a transductive inference algorithm to alleviate the low-data problem. It utilizes the unlabeled query samples to enrich the class features to be more representative. Extensive experiments show that our method is far simpler and more efficient than recent few-shot meta-learning approaches, and produces state-of-the-art results.

Acknowledgement This work is partially supported by National Key R&D Program of China (No.2017YFA0700800), Natural Science Foundation of China (NSFC): 61876171 and Beijing Natural Science Foundation under Grant L182054.

References