Learning to Prompt for Continual Learning
Zifeng Wang, Zizhao Zhang, Chen-Yu Lee, Han Zhang, Ruoxi Sun, Xiaoqi Ren, Guolong Su, Vincent Perot, Jennifer Dy, Tomas Pfister
Introduction
Contrary to ordinary supervised learning that trains on independent and identically distributed (i.i.d.) data, continual learning tackles the problem of training a single model on non-stationary data distributions where different classification tasks are presented sequentially. However, since the model only has access to the current data in an individual phase of the learning cycle, it is prone to overfit on the currently available data and suffers from performance deterioration on the previously trained data due to catastrophic forgetting .
A major body of work in continual learning follows the learning paradigm by adapting the entire or partial model weights continually as the data distribution shifts, with a focus on preserving past knowledge . Although many types of methods attain good results, there are still critical limitations that need to be addressed. First, motivated by the episodic memory in the hippocampus according to the Complementary Learning Systems (CLS) theory , many state-of-the-art methods rely on a rehearsal buffer to re-train a portion of past examples. However, they suffer from substantial performance deterioration with smaller buffer size and become ineffective when a rehearsal buffer is not allowed – for example, in real-world scenarios where data privacy matters . This suggests that simply buffering past data and re-train the model may not be the best approach to retrieve past knowledge. Without accessing a rehearsal buffer, another branch of works bypass the forgetting issue by assuming known task identity at test time, so that they are able to attach task-independent modules to the shared model for inference. However, knowing task identity at test time restricts practical usage.
The limitations of prior work bring up critical questions in continual learning : (1) Whether the form of episodic memory can go beyond buffering past data to more intelligent and succinct episodic memory system? (2) How to automatically select relevant knowledge component for arbitrary sample without knowing its task identity?
To answer the first question, we draw inspiration from recent advances in prompt-based learning (prompting) , a new transfer learning technique in the field of natural language processing (NLP). Prompting techniques design model textual inputs with templated or learnable prompt tokens containing additional task-specific information, such that the pre-trained language model can process parameterized inputs in order to perform prompt-specific prediction . Intuitively, prompt-based learning reformulates learning downstream tasks from directly adapting model weights to designing prompts that “instruct” the model to perform tasks conditionally. A prompt encodes task-specific knowledge and has the ability to utilize pre-trained frozen models more effectively than ordinary fine-tuning . Thus, it is promising to leverage prompts to learn knowledge, and further store learned knowledge, in the continual learning context.
Nevertheless, it is not clear how to apply prompting to address the aforementioned second question in continual learning directly: On one hand, if we train different prompts for different tasks in the continual learning context, test-time task identity is still required for making predictions using an appropriate task-specific prompt. On the other hand, as a transfer learning technique, the target of prompting is to make frozen pre-trained models achieve good performance on down-streaming individually, not sequentially. Therefore, if we instead maintain a single shared prompt for all tasks, the problem of catastrophic forgetting may still exist (see Section 5.4).
To this end, we propose a new continual learning method called Learning to Prompt for Continual Learning (L2P), which is orthogonal to popular rehearsal-based methods and applicable to practical continual learning scenarios without known task identity or boundaries. Figure 1 gives an overview of our method in contrast to typical continual learning methods. L2P leverages the representative features from pre-trained models; however, instead of tuning the parameters during the continual learning process, L2P keeps the pre-trained model untouched, and instead learns a set of prompts that dynamically instruct models to solve corresponding tasks. Specifically, the prompts are structured in a key-value shared memory space called the prompt pool, and we design a query mechanism to dynamically lookup a subset of task-relevant prompts based on the instance-wise input features. The prompt pool, which is optimized jointly with the supervised loss, ensures that shared prompts encode shared knowledge for knowledge transfer, and unshared prompts encode task-specific knowledge that help maintain model plasticity. Our design explicitly decouples shared and task-specific knowledge, thus largely reducing the interference between task-specific knowledge during optimization, leading to minimal catastrophic forgetting without the necessity of a rehearsal buffer. The instance-wise query mechanism removes the necessity of knowing the task identity or boundaries, enabling the most challenging, yet under-investigated task-agnostic continual learning. The selected prompts are then prepended to the input embeddings (Figure 2), which implicitly add task-relevant instruction to pre-trained models, so that the model recalls the most relevant features to conduct corresponding tasks. In summary, this work makes the following contributions:
We propose L2P, a novel continual learning framework based on prompts for continual learning, providing a new mechanism to tackle continual learning challenges through learning a prompt pool memory space, which are served as parameterized “instructions” for pre-trained models to learn tasks sequentially. The method is applicable to handle the most challenging task-agnostic continual learning.
We conduct comprehensive experiments to demonstrate the effectiveness of L2P on multiple continual learning benchmarks, including class- and domain-incremental, and task-agnostic settings. The proposed L2P outperforms previous state-of-the-art methods consistently on all benchmarks. Surprisingly, even when a rehearsal buffer is not used, L2P still achieves competitive results against rehearsal-based methods, which is ideal in real-world scenarios when rehearsal buffer is prohibited.
To the best of our knowledge, we are the first to introduce the idea of prompting in the field of continual learning. We expect that our method provides a different perspective for solving frontier challenges in continual learning.
Related Work
Here we draw connections and discuss differences between our method to related works.
Continual learning. There are three main categories of recent continual learning algorithms. Regularization-based methods limit the plasticity of the model by limiting the learning rate on important parameters for previous tasks. Although these methods address catastrophic forgetting to some extent without storing past examples, they cannot get satisfactory performance under challenging settings or complex datasets .
Rehearsal-based methods construct a data buffer to save samples from older tasks to train with data from the current task. Based on this simple yet effective idea, many recent methods improve upon it by involving additional knowledge distillation penalties , or leveraging self-supervised learning techniques . Albeit its simplicity in concept, rehearsal-based methods achieve state-of-the-art performance on various benchmarks . However, the performance of rehearsal-based methods generally deteriorates with smaller buffer size , and rehearsal-based methods are eventually not applicable to scenarios where data privacy should be taken into account . Different from directly saving data from past knowledge to re-train the model, our method stores past knowledge in small learnable prompt parameters to instruct the model to deal with current task, and in turn accumulate current knowledge to the prompts. Our method does not need a rehearsal buffer to achieve performance close to rehearsal-based methods, and could be further improved to set a new stat of the art given a small rehearsal buffer.
Architecture-based methods aim at having separate components for each task. The task-specific components can be identified by expanding the network , or attend to task-specific sub-networks . However, most methods, which require task identity to condition the network at test-time, are not applicable to more realistic class-incremental and task-agnostic settings when task identity is unknown. Some recent methods either infer task identity directly , or additionally add rehearsal buffer to bypass the problem . Nevertheless, these methods require substantial amount of additional parameters, sometimes close to the size of the full model . On the contrary, L2P does not require test-time task identity and only adds negligible amount of additional parameters (). Although L2P also introduces additional prompt parameters, it has a totally different design principle from architecture based methods: L2P designs a novel prompt-based memory to learn high-level instructions from model inputs to steer model outputs and keeps the learned architecture fixed. In contrast, most architecture-based methods aim to separate model parameters.
Lastly, recent work of CTN and DualNet start to consider knowledge management via a controller that models task-level information in addition to a backbone model. However, CTN still requires task identity at test time, while DualNet needs a rehearsal buffer to work. Moreover, CTN and DualNet are inspired from a different perspective of CLS, which suggests human beings achieve continual learning through two systems that facilitate fast learning and long-term remembering respectively. Interestingly, though we get our inspiration differently, L2P could be interpreted through CLS theory exactly: The prompt pool deals with fast learning, and the backbone model serves as long-term memory.
Prompting for transfer learning. The high-level idea of prompting is to apply a function to modify the input text, so that the language model gets additional information about the task. However, the design of a prompting function is challenging and requires heuristics. Recent work, including prompt tuning and prefix tuning , seek to address this problem by applying learnable prompts in a continuous space, achieving excellent performance on transfer learning. Prompts capture task-specific knowledge with much smaller additional parameters, than its competitors, such as Adapter and LoRA . The central idea of prompting is mainly designed for transfer learning. Note that it is non-trivial to directly apply prompting in continual learning. Our proposed novel framework reveals its values to continual learning problems.
Prerequisites
Continual learning is usually defined as training machine learning models on non-stationary data from sequential tasks. We define a sequence of tasks , where the -th task contains tuples of the input sample and its corresponding label . The goal is to train a single model parameterized by , such that it predicts the label given an unseen test sample from arbitrary tasks. Data from the previous tasks may not be seen anymore when training future tasks.
Depending on the task transition environment, continual learning can be categorized into multiple settings with slightly different challenges. The common task-, class-, and domain-incremental setting assume task data arrives in sequence in a discrete manner. Different from class-incremental, task-incremental learning assumes task identity is known at test time and are often regarded as the simplest setting . Different from task- and class-incremental settings where each task has different classes, domain-incremental learning maintains the same set of classes for every task and only changes the distribution of by task. In the more challenging task-agnostic setting, task data in changes smoothly, and the task identity is unknown. Our paper tackles the more challenging class-incremental and domain-incremental, and further explores the task-agnostic settings.
2 Prompt-based learning and baselines
Prompt-based learning is an emerging technique in NLP. In contrast to traditional supervised fine-tuning, this type of methods design task-specific prompt functions to instruct pre-trained models perform corresponding tasks conditionally . One of recent techniques, Prompt Tuning (PT) , proposes to simply condition frozen T5-like language models to perform down-stream NLP tasks by learning prompt parameters that are prepended to the input tokens to instruct the model prediction. Without loss of generality, here we introduce the definition of PT using the image modality transformer-based sequence models . The definition is easy to generalize to other modalities and sequence-based models.
Compared with ordinary fine-tuning, literature shows that prompt-based learning results in a sequence-based model having higher capacity to learn features . Despite its successes in transfer learning to train individual prompts for each task, prompting can not be directly applied to continual learning scenarios where test-time task identity is unknown.
Learning to Prompt (L2P)
The motivations of introducing prompt pool are threefold. First, the task identity at test time is unknown so training task-independent prompts is not feasible. Second, even if the task-independent prompt can be known at test time, it prevents possible knowledge sharing between similar tasks . Third, while the naive way of learning a single shared prompt for all tasks enables knowledge sharing, it still causes severe forgetting issue (see Section 5.4). Ideally one would learn a model that is able to share knowledge when tasks are similar, while maintaining knowledge independent otherwise. Thus, we propose using a prompt pool to store encoded knowledge, which can be flexibly grouped as an input to the model. The prompt pool is defined as
where ; represents concatenation along the token length dimension. Prompts are free to compose, so they can jointly encode knowledge (e.g. visual features or task information) for the model to process. Ideally, we want to achieve a more fine-grained knowledge sharing scheme via prompt combinations at the instance-wise level: similar inputs tend to share more common prompts, and vice versa.
2 Instance-wise prompt query
We design a key-value pair based query strategy to dynamically select suitable prompts for different inputs (see Figure 2). This key-valued memory query mechanism shares some design principles with methods in other fields, such as Differentiable Neural Computer and VQ-VAE , which have external memory to maintain, and employ them for a different purpose.
where represents the a subset of top- keys selected specifically for from . Note that the design of this key-value strategy decouples the query mechanism learning and prompt learning processes, which has been experimentally shown to be critical (see Section 5.4). Furthermore, querying prompts is done in an instance-wise fashion, which makes the whole framework task-agnostic, meaning that the method works without needing clear task boundaries during training, nor task identity at test time.
Optionally diversifying prompt-selection. Although our method does not need task boundary information, in real-world scenarios and experimental datasets, it is quite common that the task transition is discrete and so task boundaries are known at train time. We find that adding such a prior into our framework can help the model learn better task-specific prompts, especially when tasks have high diversity. To this end, we propose a simple extension to add task boundary prior, which is optional for L2P.
During training of task , we maintain a prompt frequency table , where each entry represents the normalized frequency of prompt being selected up until task . To encourage the query mechanism to select diverse prompts, we modify equation 3 to
where penalizes the frequently-used prompts being selected to encourage diversified selection. Equation 4 is only applicable during training; at test time, equation 3 is used.
3 Optimization objective for L2P
At every training step, after selecting prompts following the aforementioned query strategy, the adapted embedding feature is fed into the rest of the pre-trained model and the final classifier parametrized by . Overall, we seek to minimize the end-to-end training loss function:
where f_{r}^{\text{avg}}=\text{AvgPool}(f_{r}({\bm{x}}_{p})[{\color[rgb]{0,0,0}0:NL_{p}},:]), i.e., the output hidden vectors corresponding to the prompt locations are averaged before the classification head. The first term is the softmax cross-entropy loss, the second term is a surrogate loss to pull selected keys closer to corresponding query features. is a scalar to weight the loss.
Experiments
To evaluate the proposed L2P, we closely follow the settings proposed in prior works , and conduct comprehensive experiments. In particular, we mainly consider (1) the class-incremental setting, where the task identity is unknown during inference; (2) the domain-incremental setting, where the input domain shifts over time; (3) the task-agnostic setting, where there is no clear task boundary. We carefully compare L2P with state-of-the-art (SOTA) methods of different categories under proper experiment settings. Moreover, we conduct extensive ablation studies to provide a deeper understanding of our method.
We compare L2P against several baselines and state-of-the-art (SOTA) continual learning methods. Our method is based on a pre-trained ViT-B/16 , which has become a common asset in advanced vision communities. We carefully choose compared methods in the same environment for fair comparison. Many recent methods claimed SOTA performance in the simplest task-incremental setting, where task identity is known at test time . We do not include these methods, since they are not applicable to more general class-incremental setting. We refer to multiple recent reviews papers and recent work and select the most well-recognized and best-performing methods. For completeness, we also include naive sequential training approaches and representative regularization-based methods. Moreover, we refer to the original codebases for implementation and hyperparameter selection to ensure the best possible performance.
Baseline methods. Upper-bound is the usual supervised finetuning on the i.i.d. data of all tasks, which is the usually regarded as the upper bound performance a method can achieve. FT-seq-frozen is the naive sequential fine-tuning approach with the pre-trained model frozen. FT-seq instead fine-tunes pre-trained model weights as well. EWC and LwF are representative regularization-based approaches that are widely compared.
SOTA rehearsal-based methods. We select 5 advanced rehearsal-based methods to compare, including ER , GDumb , BiC , DER++ and Co2L . ER and GDumb are simple in concept, but they have achieved very strong performance not only in their own work, but in later literature as well. DER++ and Co2L are the latest SOTA methods.
SOTA architeture-based methods. We select two representative architecture-based methods to compare. SupSup and DualNet are both based on ResNet18, recommended by their original authors. We compare the relative performance to the corresponding upper-bound performance for fairness.
Our methods. L2P is our proposed method without rehearsal buffer. L2P-R is L2P equipped with a rehearsal buffer for a fair comparison with SOTA methods.
2 Datasets and experimental details
Datasets. We use Split CIFAR-100 and 5-datasets for class-incremental setting, CORe50 for domain-incremental setting, and Gaussian scheduled CIFAR-100 for task-agnostic setting, to evaluate the effectiveness of our method. Details of the datasets are introduced in Appendix Dataset details and licensing information.
Evaluation metrics. For settings with task boundaries and where each task has an associated test set, we use two metrics, Average accuracy (higher is better) and Forgetting (lower is better), which are widely used in previous works . For settings without task boundary or where there is only a single test set available, we report the final test accuracy following the common protocol .
Training details. For L2P, we train all models using Adam with and , a batch size of 128, and a constant learning rate of for all settings. Input images are resized to and normalized to the range of $M=10,N=5,L_{p}=5M=20,N=4,L_{p}=546,08092,1600.05\%0.11\%\lambda\lambda=0.5$ consistently for all datasets. Main experimental results are averaged over 3 runs, and corresponding standard deviation is reported as well.
3 Main results
Results on class-incremental learning. Table 1 summarizes the results on these two class-incremental benchmarks. L2P outperforms all comparing methods consistently under different configurations, in terms of both average accuracy and forgetting. We observe that when the buffer size is relatively large, L2P not only outperforms all other methods, but also closes a significant part of the gap to the upper bound performance under the i.i.d. setting. When the buffer size gets smaller, L2P outperforms others by a even larger margin. Finally, when there is no buffer, rehearsal-based methods are no longer capable, while L2P still remains superior performance by beating regularization-based methods, and outperforms almost all rehearsal-based methods when buffer is small.
Table 2 shows the comparison between L2P and architecture-based methods on Split CIFAR-100. Instead of absolute performance in average accuracy, we use difference to upper-bound (Diff) to measure the performance of each method given a specific architecture. We observe that L2P outperforms both methods with (DualNet) or without (SupSup) rehearsal buffer, by a large margin.
The outstanding performance of L2P over all competing methods indicates that our proposed prompt pool successfully accumulates knowledge from experiences, thus it can overall improve the learning performance while mitigating catastrophic forgetting even without a rehearsal buffer.
Results on domain-incremental learning. Table 3 summarizes the results on the domain-incremental setting. L2P remains the best performance compared with other methods. Interestingly, all rehearsal-based comparing methods perform quite closely (except GDumb). The observation of relatively modest performance gap between baseline methods and the upper-bound result has also been reported in , thus there is indeed a significant performance gap between our method and others.
Results on task-agnostic learning. Although task-agnostic setting is usually considered more challenging , the topic is under-investigated. We conduct more exploratory studies on task-agnostic settings. Table 4 summarizes the results on the challenging task-agnostic learning setting. We do not compare with LwF, BiC and Co2L since they require task boundary to save model snapshots and calculate distillation loss. It is beyond our scope to extend them to this setting. We also use the online version of EWC proposed by for the task agnostic setting. Since all compared methods are based on pre-trained models, the absolute numbers are not too far away from Upper-bound. As can been seen, rehearsal-based methods have clear advantages. Nevertheless, L2P still achieves the best performance even when buffer size is zero, among all methods, including ones have a rehearsal buffer. We believe that the smoother transition of tasks implicitly help L2P consolidate knowledge into prompts. Since we have better prompts, the benefit of a rehearsal buffer is naturally weakened.
4 Effectiveness of core designs
Effect of prompt related components for L2P. Table 5 (row 1) removes the prompt pool design and uses a single prompt to train sequentially. The performance has a significant drop, suggesting that a single prompt suffers severe catastrophic forgetting and knowledge interference between tasks, while our design of prompt pool encodes task-invariant and task-specific knowledge well. Table 5 (row 2) removes the learnable key associated with prompts and directly uses mean of prompts as keys. As results show, learnable keys play an important role to decouple the query and prompt learning processes. Table 5 (row 3) removes the diversified prompt selection (only used in 5-dataset experiments). Basically, removing it allows instances from different tasks to choose prompts freely. The decrease in performance suggests that, when tasks are diverse, adding this strategy indeed reduces unnecessary knowledge sharing and thus mitigating interference between unrelated tasks.
To better understand the prompt selection mechanism, we plot the prompt selection histograms for each task in both Split CIFAR-100 and 5-datasets in Figure 3 under the best-performing parameters settings, respectively. From the plot of Split CIFAR-100 (left), the tasks largely share all prompts, meaning that our prompt selection mechanism encourages more knowledge sharing between similar tasks. In contrast, in the plot of 5-datasets (right), diverse tasks require more task-specific prompts and share less.
Effect of hyperparameters for L2P. Recall that there are three key hyperparameters, including the size of the prompt pool , length of a single prompt , and the selection size used as model input. Intuitively, decides the total capacity of learnable prompts. decides capacity of a singe prompt (which jointly encodes certain knowledge), and decides the total size used to prepend the input. From the results on both datasets (Figure 4 (left-middle)), a too small always negatively affects results, while an oversized prompt may introduce knowledge underfitting. We hypothesize that a reasonable capacity of a single prompt is critical to encode a certain aspect of shared knowledge. Increasing the prompt pool size shows positive effect on performance as shown in Figure 4 (right) on 5-datasets while not as effective on Split CIFAR-100, suggesting a large enough pool size is needed to encode task-specific knowledge when tasks are diverse.
Conclusion
This paper presents a novel method to address some of the key challenges in continual learning with a method that can achieve strong performance without a need for rehearsal and task identity. L2P introduces prompt-based learning to continual learning and proposes a novel technique to enable a single pre-trained model to adapt to sequential tasks via a shared prompt pool, successfully mitigating the catastrophic forgetting problem. The resulting method significantly outperforms previous SOTA on several continual learning problems, including class-incremental and domain-incremental. We show our method is general enough to handle even more challenging task-agnostic settings where previous methods are incapable of.
Acknowledgments
We would like to thank Chun-Liang Li, Jeremy Martin Kubica, Sayna Ebrahimi, Stratis Ioannidis, Nan Hua, and Emmanouil Koukoumidis, for their valuable discussions.
References
Potential negative societal impact
L2P is a strong continual learning method and has great potential to be applied in various fields. However, there are some ways it could be misused. Our method takes a well-pretrained model as a backbone, thus any bias and fairness issues in the original model may be carried over during the continual learning process. We encourage any users to thoroughly check the pretrained model to mitigate any bias and fairness issues. Moreover, the method could be deployed in safety-critical applications, such as autonomous driving systems , which may present potential security issues in terms of adversarial attacks . We would recommend testing the robustness of our method in future work and design corresponding defense techniques to deal with potential security concerns.
Limitations
Although our method is demonstrated on vision models, it does not make any assumption of modalities. We leave exploration on other modalities as future work. Additionally, L2P assumes there are pre-trained sequence-based models. While they have become common assets and future directions in advanced communities, how to generalize our framework to other vision architectures (e.g. ConvNet) could be an appealing research direction.
How to achieve continual learning that can satisfy the real-world requirements is an important direction that remains challenging. For example, the task-agnostic setting is known as the most challenging setting and is very close to real-world scenarios. Although our method takes a step further towards this goal, however, the current commonly used Gaussian scheduled CIFAR-100 is synthetic and still far from realistic. Thus, we think it also requires more complex benchmarks to evaluate the ability of task-agnostic continual learning methods and push forward the advances of this real-world challenge.
Dataset details and licensing information
Split CIFAR-100 (class-incremental). This dataset splits the original CIFAR-100 into 10 tasks, 10 disjoint classes per task. Since the tasks are from a single original dataset, they share some similarities and some classes could be from the same superclass. Although CIFAR-100 is a simple image classification dataset, it remains quite challenging for continual learning studies, especially in the class-incremental setting .
5-datasets (class-incremental). We also use a challenging dataset proposed in . This dataset consists of five image classification datasets: CIFAR-10, MNIST , Fashion-MNIST , SVHN , and notMNIST . Although each dataset alone is not hard, the sequential training of them is fairly challenging even with ImageNet pre-trained models, since models are susceptible to forgetting when the tasks are diverse .
CORe50 (domain-incremental). This is a widely used dataset specifically designed for continual object recognition . It is a collection of 50 objects collected in 11 distinct domains, where 8 of them (120,000 samples) are used for training, and the rest are considered as a single test set (45,000). Methods are trained on each domain sequentially.
Gaussian scheduled CIFAR-100 (task-agnostic). The distribution of data shifts gradually throughout the learning process , the probability that a class is present in a batch follows a Gaussian distribution centered with intervals. There is no explicit task boundaries between batches, thus requiring methods to be able to implicitly adapt to non-stationary data distribution without utilizing any task-specific information during both training and inference.
CIFAR-10 and CIFAR-100 , Fashion-MNIST are licensed under the MIT license.
MNIST is licensed under the Creative Commons Attribution-Share Alike 3.0 license.
CORe50 is under the Creative Commons Attribution 4.0 International license.
The licensing information is not available for SVHN, notMNIST .
Algorithm details
To better illustrate our proposed method, we present a whole picture of the training procedure in Algorithm 1. Note that for prediction, we simply replace loss calculation to label prediction. Optionally, we can replace the top- keys lookup by equation 4, when task boundary prior is known.