Supervised Contrastive Replay: Revisiting the Nearest Class Mean Classifier in Online Class-Incremental Continual Learning

Zheda Mai, Ruiwen Li, Hyunwoo Kim, Scott Sanner

Introduction

With the ubiquity of personal smart devices and image-related applications, a massive amount of image data is generated daily. A practical online learning system is expected to learn incrementally without storing all streaming data and retraining over it due to space and computational resource limitations. However, a well-documented drawback of deep neural networks that prevents it from learning continually is called catastrophic forgetting — the inability to retain previously learned knowledge after learning new tasks. To address this challenge, Continual Learning (CL) studies the problem of learning from a non-i.i.d stream of data, intending to preserve and extend the acquired knowledge while minimizing storage, computation, and time.

Most early CL approaches considered task-incremental settings, in which new data arrives one task at a time, and the model can utilize task-IDs during both training and inference time . This setting implicitly simplifies the CL problem as the model just needs to classify labels within a task with the help of task-IDs. Meanwhile, this simplification diminishes the applicability of this setting when task-IDs are not available. In this work, we consider a more realistic but challenging setting, known as online class-incremental, where a model is required to learn new classes continually from an online data stream (each sample is seen only once) and classify all labels without task-IDs.

Current CL methods can be taxonomized into three major categories: regularization, parameter-isolation, and replay methods . The replay approach has been shown to be simple and efficient compared to other approaches in the online class-incremental setting . However, A key challenge of replay methods is the imbalance between old and new classes, as only a small amount of old class data are stored in the replay buffer. Recent works have revealed that the Softmax classifier and its associated fully-connected (FC) layer are seriously affected by the class imbalance, which leads to task-recency bias — the tendency of a model to be biased towards classes from the most recent task . Although the Nearest-Class-Mean (NCM) classifier is significantly undervalued in the CL community, we demonstrate that it is a simple yet effective substitute for the Softmax classifier as it not only addresses the recency bias but also avoids structural changes in the FC layer when new classes are observed. Moreover, we observe considerable and consistent performance gains when replacing the Softmax classifier with the NCM classifier for five methods with memory buffers. Since also observed similar gains in methods without memory buffer, we advocate using the NCM classifier instead of the commonly used Softmax classifier for future study.

Furthermore, to exploit the NCM classifier more effectively, the data embeddings belonging to the same class should be clustered and well-separated from those with different class labels. To this end, we contribute Supervised Contrastive Replay (SCR), which leverages the supervised contrastive loss to explicitly encourage samples from the same class to cluster tightly in embedding space and push those of different classes further apart when replaying buffered samples with the new samples. Through extensive experiments on three commonly used benchmarks in the CL literature, we demonstrate that SCR outperforms state-of-the-art methods by significant margins with three different memory buffer sizes.

Related Work

The model receives a small batch BtnB_{t}^{n} of size bb from task DnD_{n} at time tt. ff and gg will be updated based on BtnB_{t}^{n} and data in Mt−1M_{t-1}, a bounded memory that can be used to store a subset of the training samples or other useful data . Moreover, we adopt the single-head evaluation setup where the classifier has no access to task-IDs during inference and hence must choose among all labels. Our goal is to train the model (f,g)(f,g) to continually learn new classes from the data stream without forgetting.

Approaches

As previously discussed, current CL methods can be classified into three major categories: regularization, parameter-isolation, and replay methods . Regularization methods constraint the updates of some important network parameters to mitigate catastrophic forgetting. This is done by either incorporating additional penalty terms into the loss function or modifying the gradient of parameters during optimization . Other regularization methods imposed knowledge distillation techniques to penalize the feature drift on previous tasks . Parameter-isolation methods bypass interference by allocating different parameters to each task . Replay methods deploy a memory buffer to store a subset of data from previous tasks for replay . Regularization methods mostly protect the model’s ability to classify within a task, and thus they do not work well in our setting, which requires the ability to classify from all labels the model has seen before . Also, most parameter isolation methods require task-IDs during inference, which violates our setting. Therefore, in this work, we will focus on replay methods, which have been shown to be efficient and effective compared to other approaches in the online class-incremental setting .

Metrics

We use the average accuracy of the test sets from observed tasks to measure the overall performance . In Average Accuracy, ai,ja_{i,j} is the accuracy evaluated on the held-out test set of task jj after training the network from task 1 to ii. By the end of training all NN tasks, the average accuracy can be calculated as follows:

2 Contrastive Learning

The general goal of contrastive learning is intuitive: the representation of “similar” samples should be mapped close together in the embedding space, while that of “dissimilar” samples should be further away . When labels are not available (self-supervised), similar samples are often formed by data augmentations of the target sample while dissimilar samples are often drawn randomly from the same batch of the target sample or from the memory bank/queue that stores feature vectors . When labels are provided (supervised), similar samples are those from the same class and dissimilar samples are those from different classes . Contrastive learning has recently attracted a surge of interest and shown promising results in various areas including computer vision , natural language processing , audio processing , graph and multimodal data .

Method

Softmax classifier with cross-entropy loss has been a standard approach for classification tasks for neural networks . Although this combination also dominates the CL for image classification, it may not be the best choice for CL due to the following deficiencies.

Architecture modification for new classes When the model receives new classes, the Softmax classifier requires the model to stop training and add weights in the FC classification layer to accommodate the new classes.

Decoupled representation and classification In the class-incremental setting, as mentioned in , it is problematic that the weights in the classification layers are decoupled from the encoder since whenever the encoder changes, weights in the classification layer must also be updated.

Task-recency bias Multiple previous works have observed that a model with the Softmax classifier has a strong prediction bias towards the most recent task due to the imbalance of new and old classes, which is the primary source of catastrophic forgetting. Figure 3 (a) shows the confusion matrix after training task 10, which shows that the model tends to predict most samples as classes in the most recent task. As illustrated in Figure 4, the means of weights for the new classes in the FC layer are much higher than those for the old classes and hence the model assigns a larger probability mass for predicting a sample as a new class vs. an old class.

Nearest Class Mean (NCM) Classifier

The NCM classifier and its variants have been widely used in few-shot or zero-shot learning . Concretely, after the embedding network ff is trained, the NCM classifier computes a class mean (prototype) vector for each class using all the embeddings of this class. To predict a label for a new sample x\mathbf{x}, NCM compares the embedding of x\mathbf{x} with all the prototypes and assigns the class label with the most similar prototype: {ceqn}

Although the NCM classifier is significantly undervalued in the CL community, we argue that it is a simple yet effective substitute for the Softmax classifier as it not only resolves the deficiencies of the Softmax classifier mentioned above but also demonstrates a considerable improvement.

Since the NCM classifier simply compares the embedding of the test sample with prototypes, it does not require an additional FC layer, and therefore, new classes can be added without any architecture modification.

As the prototypes change instinctively based on the encoder, the NCM classifier is more robust against changes of the encoder.

The biased weights in the FC layer result in the task-recency bias, but since the NCM classifier does not involve the FC layer, it is intrinsically less prone to the task-recency bias.

Figure 5 shows the average accuracy comparison of a Softmax classifier and an NCM classifier on five methods. NCM classifiers show significant improvements over the commonly used Softmax classifier across all five methods and three datasets, which suggests that the dominance of the Softmax classifier in online continual learning should be revisited.

Although the NCM classifier has shown impressive results, the embedding quality greatly and directly impacts the performance of the NCM classifier. To effectively exploit the NCM classifier, the data embeddings belonging to the same class should be clustered and well-separated from those with a different class label. However, the binary cross-entropy loss used in iCaRL may not be capable of addressing the relationship between classes, and the commonly used categorical cross-entropy loss may not be effective in creating discernible patterns in the embedding space, as shown in Figure 2.

2 Supervised Contrastive Replay

where II is the set of indices of BIB_{I} and A(i)=I\{i}A(i)=I\backslash\{i\}, represents the set of indices of all samples in BIB_{I} except for sample ii. P(i)≡{p∈A(i):yp=yi}P(i)\equiv\left\{p\in A(i):{\boldsymbol{y}}_{p}={\boldsymbol{y}}_{i}\right\} is the set of all positives (i.e., samples with the same labels as sample ii) in BIB_{I} excluding sample ii, and ∣P(i)∣|P(i)| is its cardinality. ZI={zi}i∈I={Proj(Enc(xi)}i∈IZ_{I}=\{z_{i}\}_{i\in I}=\{{Proj}({Enc}({{x_{i}}})\}_{i\in I}; τ∈R+\tau\in\mathcal{R}^{+} is an adjustable temperature parameter controlling the separation of classes; the ⋅\cdot indicates the dot product.

Supervised Contrastive Replay (SCR)

An overview of SCR can be found in Figure 1. As mentioned in Section 2.1, during the training phase, the model receives one small batch BnB_{n} at a time from task DnD_{n} in the data stream D\mathcal{D}. An input batch is created by concatenating BnB_{n} with another batch BMB_{\mathcal{M}} selected from the memory buffer M\mathcal{M}. The input batch and its augmented view are encoded by a shared encoder network Enc(⋅)Enc(\cdot) and a projection network Proj(⋅)Proj(\cdot) before the representations are evaluated by the supervised contrastive loss LSCL\mathcal{L}_{\text{SCL}}. After updating both Enc(⋅)Enc(\cdot) and Proj(⋅)Proj(\cdot) with the gradient from LSCL\mathcal{L}_{\text{SCL}}, the memory buffer M\mathcal{M} will be updated with BnB_{n}.

During the testing phase, Proj(⋅)Proj(\cdot) is discarded. All the buffered samples are fed into Enc(⋅)Enc(\cdot) to obtain the embeddings, which are used to compute the class means (prototypes) for the NCM classifier. As SCR builds much more discernible patterns in the embedding space with the contrastive loss, the NCM classifier is able to unleash its capability in our method. Algorithm 1 summarizes the training and inference procedures.

Experiment

Split CIFAR-10 is constructed by splitting the CIFAR-10 dataset into 5 different tasks with non-overlapping classes and 2 classes in each task, similarly as in . Split CIFAR-100 splits the CIFAR-100 dataset into 10 disjoint tasks, and each task has 10 classes. Split Mini-ImageNet divides the Mini-ImageNet dataset into 10 disjoint tasks with 10 classes per task.

Baselines

We compare our proposed SCR against several state-of-the-art continual learning algorithms:

A-GEM (ICLR’19) : Averaged Gradient Episodic Memory, that utilizes the samples in the memory buffer to constrain the parameter updates.

ASERμ (AAAI’21) : Adversarial Shapley Value Experience Replay that leverages Shapley value adversarially in memory retrieval.

ER (ICML-W’19): Experience replay, a replay method with random sampling in memory retrieval and reservoir sampling in memory update.

EWC++ (ECCV’18) : An online version of EWC , a regularization method that limits the update of parameters that were crucial to the past tasks.

GSS (NeurIPS’19) : Gradient-Based Sample Selection, a replay method that diversifies the gradients of the samples in the replay memory.

LwF (TPAMI’18) Learning Without Forgetting, a regularization method that utilizes knowledge distillation to penalize the feature drifts on previous tasks.

MIR (NeurIPS’19) : Maximally Interfered Retrieval, a replay method that retrieves memory samples with loss increases given the estimated parameter update based on the current batch.

offline: This is not a CL method, but rather an upper bound; offline trains the model over multiple epochs on the whole dataset with iid sampled mini-batches. We use 50 epochs for offline training.

fine-tune: A lower-bound method that simply trains the model when new data is presented without any measure for forgetting avoidance.

Implementation Detail

Following , we use a reduced ResNet18 as the backbone model for all datasets. We use stochastic gradient descent with a learning rate of 0.1, and the model receives a batch with size 10 at a time from the data stream. All the methods except for SCR are trained with cross-entropy loss and classify with the Softmax classifier. The projection network of SCR is a Multi-Layer Perceptron (MLP) with one hidden layer (ReLU) and an output size 128, and we set the temperature τ\tau to 0.1. We use reservoir sampling for memory update and random sampling for memory retrieval and use a memory batch size 100. The ablation study of the variables mentioned above will be discussed in Section. 4.4.

2 Evaluation of NCM Classifier

To assess the effectiveness of the NCM classifier, we compare five methods that employ memory buffers (AGEM, ER, GSS, MIR, ASERμ) with their variants equipped with the NCM classifier. As we can see in Figure 5 and Table 1, methods with the NCM classifier show significant improvements over those with the default Softmax classifier. For instance, in CIFAR100, the NCM classifier helps ASERμ with 1k memory achieve 22%, which requires five times more memory to achieve when using the Softmax classifier. Generally, we also observe that the performance gain is more notable when the memory buffer is small. For example, in Mini-ImageNet, MIR obtains 66.4% relative improvement (10.3% →\rightarrow 17.8%) with M=1k, which is only improved by 27.7% relatively (17.3% →\rightarrow 22.1%) with M=5k. Furthermore, the NCM gains are less obvious for GSS, and we find out that it’s because some classes only have a few or sometimes zero samples in the GSS buffer, which makes it hard to estimate the correct prototypes for those classes. Moreover, ASERμ has better NCM gains in general, and it’s because ASERμ tends to learn more discernible embeddings, as we can see in Figure 1.

To sum up, we observe considerable and consistent performance gains when replacing the Softmax classifier with the NCM classifier for five methods on three different datasets and memory sizes. Since also observed similar gains in methods without memory buffer, we advocate using the NCM classifier instead of the commonly used Softmax classifier for future study.

3 Evaluation of SCR

To evaluate the performance of SCR, we compare it with several state-of-the-art CL methods described in Section 4.1. As we can see in Figure 6, SCR consistently outperforms all the compared methods by enormous margins along the whole data streams of three different datasets. Note that all the compared methods on the plots have already been NCM-augmented. Table 1 shows the detailed comparison of SCR with all the compared methods on different datasets and memory sizes. The last row of the table shows the absolute improvements over the second-best methods. SCR consistently achieves state-of-the-art results across all settings and outperforms the compared methods by large margins. SCR achieves 35.4% (13.3\%{\color[rgb]{1,0,0}\uparrow}), 37.8% (8.2\%{\color[rgb]{1,0,0}\uparrow}) and 65.7% (15.4\%{\color[rgb]{1,0,0}\uparrow}) respectively in Mini-ImageNet, CIFAR100 and CIFAR10 respectively. The success of SCR comes from (i) the NCM classifier, which has shown impressive performance over the Softmax classifier in Section 4.2, and (ii) the contrastive loss, which enables the model to learn more discernible embeddings and provides a solid foundation for the NCM classifier. Moreover, we observe SCR benefits from a large memory buffer in general, as contrastive learning desires more diverse negative samples. For example, SCR achieves 60.2% relative gain with M=5k (21.1% (MIR-NCM) →\rightarrow 35.4%), while obtains 35.4% relative improvement with M=1k (16.6% (MIR-NCM) →\rightarrow 24.1%). In terms of task-recency bias, we can see in Figure 3 (b) that SCR is clearly much less biased than ER even though a slight bias is still observed. Furthermore, SCR does not sacrifice its computation efficiency, as shown in Figure 7. Its running time (combined training and inference) is shorter than ASERμ and only slightly longer than MIR.

In summary, by evaluating on three standard CL datasets and comparing to the state-of-the-art CL methods, we have strongly demonstrated the effectiveness and efficiency of SCR in overcoming catastrophic forgetting, which brings online CL much closer to its ultimate goal of matching offline training while maintaining a low computation footprint.

4 Ablation Study

In this subsection, we aim to explore the impact of various SCR configurations on its performance. We use SCR with M=2k on CIFAR100 as the study case to analyze the impacts of components of SCR.

Figure 8 (a) shows the impact of the memory batch size. Generally speaking, contrastive learning benefits from larger batch sizes as it means more negative samples . Nevertheless, in online CL, accuracy improvement is more obvious with the increase of BMB_{\mathcal{M}} when BMB_{\mathcal{M}} is smaller than 200. The performance drops when BMB_{\mathcal{M}} continues to increase. We suspect the decrease is due to the overfitting of the memory samples as 500/1,000 are 25%/50% of the whole memory buffer in this study case.

Impact of memory buffer management.

We compare random retrieval + reservoir update(Random), ASER retrieval + update (ASER), ASER retrieval (ASER-R), ASER update (ASER-U) and GSS. As we can see in Figure 8 (b), the random option is much better than GSS and slightly better than others. We observed that some classes have only a few or zero samples in the memory for GSS, which is undesirable for SCR. Although random seems reasonable for the balanced CIFAR100 dataset, when facing imbalanced datasets, combing SCR with other memory management methods may yield better performance .

Impact of temperature variable τ𝜏\tau.

We can see from Figure 8 (c) that the performance deteriorates when the τ\tau is too low and too high. SCR with τ\tau ranging from 0.02 to 0.16 achieves stable results.

Impact of projection network P​r​o​j​(⋅)𝑃𝑟𝑜𝑗⋅Proj(\cdot).

We tried Multi-Layer Perceptron (MLP), linear and no projection network (None). Although suggests that a nonlinear projection network improves the representation quality, we find that the choice of projection network is insignificant in online CL as shown in Figure 8 (d).

Conclusion

In this paper, we first demonstrated that the NCM classifier is a simple yet effective substitute for the Softmax classifier in the online CL. It resolves several deficiencies of the Softmax classifier and shows considerable and consistent performance gains across a variety of CL methods. Based on these results, we advocate using the NCM classifier instead of the commonly used Softmax classifier for future study of CL methods. Moreover, to leverage the NCM classifier more effectively, we proposed SCR that explicitly encourages samples from the same class to cluster tightly in embedding space while pushing samples of different classes further apart during experience replay-based training.

Empirically, we observe that our proposed SCR substantially reduces catastrophic forgetting in comparison to state-of-the-art CL methods and outperforms them all by a significant margin on various datasets and memory settings. In summary, leveraging a simple randomized experience replay method while using a supervised contrastive loss (in place of cross-entropy) combined with an NCM classifier bring us closer to realizing the ultimate goal of continual learning to perform as well as offline training methods.

References