Learning to Prompt for Continual Learning

链接:[2112.08654] Learning to Prompt for Continual Learning (arxiv.org)

代码:https://github.com/google-research/l2p

来源:CVPR2022 

1.背景

灾难性遗忘的两种解决办法:

1.存储旧任务的少类样本(rehearsal)(缓冲区较小的情况下,性能大幅下降)

2.获取task identity,将类增量变为任务增量(限制了实际使用)

解决思路:Prompt

Prompt使用包含额外任务特定信息的模板化或可学习提示标记设计文本输入模型,这样预训练的语言模型可以处理参数化输入,以便执行特定提示的预测。

Prompt将学习下游任务从直接调整模型权重改为设计提示“指导”模型有条件地执行任务。提示编码特定于任务的知识,比普通微调更有效地利用预训练的冻结模型。

prompt learning:可以使基于序列的模型具有更高的学习特征的能力

instance-wise:样例水平

本文提出了L2P:使用单个主干模型,并学习一个提示池来有条件地指示模型。

2.模型

2.1思想:

特定于任务的知识存储在prompt pool中;以实例方式(instance-wise)自动选择和更新池中的提示,因此在测试时不需要任务标识。prompt space<224X224

domain-incremental:为每个任务维护相同的类集,并且只更改x按任务的分布

task agnostic:数据变化平稳,任务标识t在训练时未知

prompt tuning基本思想:

预先设定可学习的参数P_{e}\in{\mathbb{R}}^{L_{p}\times D}(称为提示符)调用嵌入特性{{x}}_{p}=[P_{e};{{x}}_{e}],并将扩展序列提供给模型函数f_{r}({{x}}_{p})以执行分类任务。{{x}}_{e}=f_{{}_{e}}(x)\in{\mathbb{R}}^{L\times D}是patch image经过embedding的结果。

但prompt不能直接用于CL,需要改造成prompt pool。

2.2模型:


2.2.1.引入prompt pool的动机:

1).测试时的task identity未知,因此训练独立于任务的prompt不可行

2).即使能学到独立于任务的prompt,它会限制类似任务之间可能的知识共享。

3).在所有任务上学习单一共享prompt的方式能够实现知识共享,但会导致严重遗忘

prompt pool池:\mathbf{P}=\{P_{1},P_{2},\cdots,P_{M}\},\quad\text{$M=$ total $\#$ of prompts},

其中P_{j}\in{\mathbb{R}}^{L_{p}\times D}的token长度为Lp,embedding size D与Xe相同。

{{x}}_{p}=[P_{s_{1}};\cdots;P_{s_{N}};{{x}}_{e}],\quad 1\leq N\leq M,;表示沿token长度维度的连接

prompt 可以自由组合,因此它们可以联合编码模型要处理的知识(例如视觉特征或任务信息)

此外,x 不需要task index t,因此适用于task agnostic。

2.2.2.Instance-wise prompt query

如何为不同的输入动态选择合适的提示?基于KEY-VALUE对的查询策略

将每个提示作为VALUE关联到一个可学习KEY:\{({{k}}_{1},P_{1}),({{k}}_{2},P_{2}),\cdots,({{k}}_{M},P_{M})\}{{k}}_{i}\in{\mathbb{R}}^{D_{k}}

理想情况下,这种query应该是基于input image本身的。

引入query function:q:{\mathbb{R}}^{H\times W\times C}\to{\mathbb{R}}^{D_{k}},将输入x编码为与key相同的维度

q({{x}})=f({{x}})[0,:](使用[class]对应的特征向量),设计原则:对于不同的任务,q是一个确定性函数,且没有可学习的参数。

{\mathbf{K}_{{x}}}=\underset{\{s_{i}\}_{i=1}^{N}\subseteq%
[1,M]}{\operatorname{argmin}}\quad\sum_{i=1}^{N}\gamma\left({q({{x}}),{{%
k}}_{s_{i}}}\right)

使用余弦距离进行query和key的匹配,得到 topN个key,且该键值策略的设计将查询机制学习和提示学习过程解耦。

这种query方式以实例方式进行,训练时不需要任务边界,测试时不需要任务标识。

2.2.3.task boundary information

虽然不是必须的,但添加这样的先验知识可以帮助模型更好地学习特定于任务的提示。

在task t train期间:维护prompt频率表H_{t}=[h_{1},h_{2},\cdots,h_{M}],每个条目代表在任务t之前选择的prompt Pi的标准化频率,鼓励diversified selection

2.3.loss

Loss function

第一项是softmax交叉熵损失,第二项是surrogate loss损失,用于将选定的键拉近到相应的查询特征。

3.结果

setting :Upper-bound:在所有数据上supervised finetuning

Upper-bound:在所有数据上supervised finetuning

消融:

最后编辑于
©著作权归作者所有,转载或内容合作请联系作者
【社区内容提示】社区部分内容疑似由AI辅助生成,浏览时请结合常识与多方信息审慎甄别。
平台声明:文章内容(如有图片或视频亦包括在内)由作者上传并发布,文章内容仅代表作者本人观点,简书系信息发布平台,仅提供信息存储服务。

相关阅读更多精彩内容

友情链接更多精彩内容