跳转至

Learning Probabilistic Prompt for Continual Learning

会议: ECCV 2026
论文: ECCV 原文
项目: http://cvlab.yonsei.ac.kr/projects/ProbPrompt
领域: 持续学习 / 提示学习 (Continual Learning)
关键词: Continual Learning, Class-Incremental Learning, Prompt Tuning, Probabilistic Prompt, Prompt Collapse

一句话总结

针对基于提示的类增量持续学习中固定确定性提示极易趋同的“提示坍塌”(Prompt Collapse)难题,本文提出将提示建模为高斯概率分布,根据输入图像查询特征动态构建混合分布并随机采样加权,配合连续步间的分布正则化 KL 散度约束,在不增加额外网络分支的前提下显著提升提示多样性并跨越式提升多基准准确率。

研究背景与动机

深度神经网络在连续学习互斥类别的新任务时,通常面临严重的灾难性遗忘问题。传统方案依赖参数正则化或外部回放缓冲区,但在大规模 Vision Transformer(ViT)架构下,全参数微调或重放不仅计算与显存开销巨大,遗忘倾向也更为剧烈。近年来基于提示的持续学习(如 L2P、DualPrompt、CODA-P、VQ-Prompt 等)成为主流范式:它们通过冻结预训练骨干网络以锁死既有通用表征,仅引入少量与输入查询特征相关联的可学习提示向量(Prompt Tokens),通过前缀微调(Prefix Tuning)使模型适应新任务。

然而,现有提示方法的核心假定是静态确定性嵌入能够充分表征图像特征分布。在类增量持续学习场景下,单个任务内部往往因包含多种差异巨大的类别而具有极高的任务内方差(如 ImageNet-R 中素描、涂鸦、卡通等多元风格混合),随着连续增量学习的推进,不同任务和类别间的数据分布更趋复杂。本文深入分析发现,现存确定性提示方法普遍陷入严重的“提示坍塌”(Prompt Collapse)陷阱——不同提示向量之间的两两余弦相似度极高(平均相似度常高达 0.88–0.96),提示之间高度相关退化,丧失了表达多样化特征分布的能力;即使盲目扩充提示池容量,也无法从根本上打破向量同质化。

解决该瓶颈的根本出路在于打破确定性向量表征的固有表达上限,将提示学习拓展到概率建模空间。核心 idea:将提示池中的每个提示分量显式建模为对角高斯分布,依据输入图像的查询特征自适应融合出混合分布并执行重参数化随机采样,再通过查询相似度重加权聚合为单个提示送入冻结 ViT,并辅以时序分布正则化损失锁定模型稳固性。

方法详解

整体框架

ProbPrompt 的整体工作流如图所示。系统冻结预训练 ViT 骨干网络,输入图像先提取用于匹配的查询特征 \(q\)。对于需要生成的 \(M\) 个提示标记,每个提示标记对应一个包含 \(N\) 个高斯分布分量的提示池。算法首先基于马氏距离度量查询特征与各高斯分量的相关度权重,构建自适应高斯混合分布;接着利用重参数化技巧从混合分布中并行抽取 \(N_s\) 个候选样本,并依据每个样本与查询特征的余弦相似度做自适应加权平均,融合成最终的单个提示标记;最后拼接 \(M\) 个提示并拼接到 ViT 输入序列前缀中执行单次前向推理。在训练过程中,采用任务交叉熵分类损失与分布正则化损失联合端到端优化所有分布的均值和方差。

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    A["输入图像与查询提取<br/>冻结 ViT 提取特征 q"] --> B["自适应高斯混合分布构建<br/>基于马氏距离计算分量权重"]
    B --> C["随机提示采样与查询重加权聚合<br/>重参数化采样并按余弦相似度融合"]
    C --> D["前缀拼接与单次前向推理<br/>Prefix-Tuning 注入冻结 ViT"]
    D --> E["双重目标协同优化<br/>交叉熵分类损失 + 跨步 KL 分布正则化"]

关键设计

1. 提示概率化建模与自适应高斯混合分布:从根本上打破提示坍塌

针对确定性向量缺乏多样性表达力、导致特征坍塌的问题,本文将提示池中的第 \(n\) 个提示分量建模为 \(D\) 维多元高斯分布 \(\mathcal{N}(\mu(m,n), \Sigma(m,n))\),其中 \(\Sigma(m,n) = \text{diag}(\sigma^2(m,n))\) 为对角协方差矩阵(实际实现中参数化为对数标准差以保障数值稳定)。给定输入图像的查询特征 \(q \in \mathbb{R}^D\),模型计算 \(q\) 与第 \(n\) 个高斯分量之间的马氏距离:

\[S^2(q, \mu(n), \Sigma(n)) = (q - \mu(n))^\top \Sigma(n)^{-1} (q - \mu(n))\]

由于协方差矩阵为对角阵,该距离可沿通道维度高效逐元素计算。基于该距离,通过 Softmax 归一化得到该分布对当前查询的相关度得分 \(s(n)\)

\[s(n) = \frac{\exp(-S^2(q, \mu(n), \Sigma(n)))}{\sum_{n'=1}^N \exp(-S^2(q, \mu(n'), \Sigma(n')))}\]

随后将这 \(N\) 个基底分布融合成一个紧凑的等效高斯混合分布 \(\mathcal{N}_{\text{GM}}(\mu_{\text{GM}}, \Sigma_{\text{GM}})\),其混合均值与方差显式表达为:

\[\mu_{\text{GM}} = \sum_{n=1}^N s(n) \mu(n), \quad \sigma_{\text{GM}}^2 = \sum_{n=1}^N s(n) \left(\sigma^2(n) + \mu(n)^2\right) - \mu_{\text{GM}}^2\]

这赋予了模型在连续空间中探索不同样本表征的能力,消除了确定性聚类带来的强硬割裂与模式坍塌。

2. 随机采样与查询自适应重加权聚合:单次前向保证高效率与语义聚焦

如果仅从混合分布中进行单次随机采样,随机噪声会导致训练过程剧烈震荡且难以收敛;若在网络深层进行多路 Monte Carlo 采样前向传播,计算和显存开销将成倍激增。为此,ProbPrompt 设计了“输入前采样、依查询加权平均”的轻量化聚合机制。通过重参数化技巧从混合分布中采样 \(N_s\) 个候选提示:

\[\tilde{P}_k = \mu_{\text{GM}} + \sigma_{\text{GM}} \odot \epsilon_k, \quad \epsilon_k \sim \mathcal{N}(0, \mathbf{I})\]

若将这 \(N_s\) 个候选向量直接进行简单平均(\(w_k = 1/N_s\)),无差别的平权操作会稀释与当前样本强相关的特异性模式。因此,本文依据每个采样候选 \(\tilde{P}_k\) 与查询特征 \(q\) 之间的余弦相似度计算 Softmax 权重 \(w_k\)

\[w_k = \frac{\exp(\text{sim}(\tilde{P}_k, q))}{\sum_{j=1}^{N_s} \exp(\text{sim}(\tilde{P}_j, q))}, \quad \hat{P} = \sum_{k=1}^{N_s} w_k \tilde{P}_k\]

拼接所有 \(M\) 个提示标记得到 \(\hat{\mathbf{P}} = [\hat{P}(1), \dots, \hat{P}(M)] \in \mathbb{R}^{M \times D}\),再统一通过 Prefix Tuning 注入 Transformer 编码器。整个采样与加权操作完全在线性时间复杂度内完成,进入 ViT 时已压缩为确定性的单个提示前缀,因此整个网络推理仅需单次前向传播,显存与耗时几乎与标准确定性提示方法完全持平。

3. 时序分布正则化损失:平抑分布漂移以锁死既有知识

在类增量持续学习中,直接引入概率采样极易因为新任务梯度的频繁扰动而导致先前任务建立的高斯分布发生剧烈形变(即模型的可塑性破坏了稳固性)。为保证提示分布在连续任务间的平滑演进并保留旧任务记忆,本文提出了跨步分布正则化损失 \(\mathcal{L}_{\text{DR}}\)。在连续训练步之间缓存前序步的分布参数 \(\mathcal{N}(\hat{\mu}(m,n), \hat{\Sigma}(m,n))\),并通过最小化当前分布与缓存分布之间的 KL 散度对其施加刚性约束:

\[\mathcal{L}_{\text{DR}} = \frac{1}{MN} \sum_{m=1}^M \sum_{n=1}^N \text{KL}\left(\mathcal{N}(\mu(m,n), \Sigma(m,n)) \,\parallel\, \mathcal{N}(\hat{\mu}(m,n), \hat{\Sigma}(m,n))\right)\]

总损失函数由任务分类交叉熵损失与加权分布正则化项构成:\(\mathcal{L} = \mathcal{L}_{\text{CE}} + \lambda \mathcal{L}_{\text{DR}}\)。该项强制提示空间在概率意义上平滑更新,为模型提供了持续学习急需的记忆锚点。

损失函数 / 训练策略

骨干网络采用 ViT-B/16 且全程冻结,仅可学习提示分布参数(均值 \(\mu\) 与对数标准差 \(\log \sigma\))以及分类器头参与反向传播。优化器采用 AdamW(\(\beta_1 = 0.9, \beta_2 = 0.999\)),初始学习率设为 \(2.5 \times 10^{-3}\),配合余弦退火策略。各基准统一训练 20 个 epoch,批次大小在 CIFAR-100/CUB-200 上为 128,在 ImageNet-R 上为 64。平衡系数 \(\lambda\) 设为 \(10^{-6}\),采样数 \(N_s = 30\),提示池参数设定为 \(M=8, N=10\)

实验关键数据

主实验

在代表性增量学习基准 ImageNet-R(测试域偏移与风格多变性)、CIFAR-100(通用目标分类)以及 CUB-200(细粒度鸟类识别)上,评估最终平均准确率(FAA)和累积平均准确率(CAA)。所有实验均基于 5 次随机种子运行的均值与标准差。

Table 1: ImageNet-R 上的 5、10、20 任务增量结果对比(ViT-B/16 预训练权重):

方法 5-Task FAA (%) 5-Task CAA (%) 10-Task FAA (%) 10-Task CAA (%) 20-Task FAA (%) 20-Task CAA (%)
FT(全微调) 18.74 ± 0.44 48.39 ± 0.58 10.12 ± 0.51 35.23 ± 0.92 4.75 ± 0.40 22.80 ± 0.37
FT++ 60.42 ± 0.87 71.59 ± 0.50 48.93 ± 1.15 66.79 ± 0.92 35.98 ± 1.38 59.68 ± 0.95
L2P (CVPR'22) 70.83 ± 0.58 78.34 ± 0.47 69.29 ± 0.73 78.30 ± 0.69 65.89 ± 1.30 77.15 ± 0.65
DualPrompt (ECCV'22) 73.05 ± 0.50 79.47 ± 0.40 71.32 ± 0.62 78.94 ± 0.72 67.87 ± 1.39 77.42 ± 0.80
CODA-P (CVPR'23) 76.51 ± 0.38 82.04 ± 0.54 75.45 ± 0.56 81.59 ± 0.82 72.37 ± 1.19 79.88 ± 1.06
HiDePrompt (NeurIPS'23) 76.29 ± 0.10 78.77 ± 0.11 76.74 ± 0.18 78.76 ± 0.11 76.46 ± 0.06 78.76 ± 0.11
EvoPrompt (AAAI'24) 77.16 ± 0.18 82.22 ± 0.54 76.83 ± 0.08 82.09 ± 0.68 74.41 ± 0.23 80.96 ± 1.42
VQ-Prompt (NeurIPS'24) 79.23 ± 0.29 82.96 ± 0.50 78.71 ± 0.22 83.24 ± 0.68 78.10 ± 0.22 82.70 ± 1.16
APT (ICCV'25) 79.20 ± 0.38 83.07 ± 0.45 79.05 ± 0.41 83.41 ± 0.54 75.94 ± 0.04 79.46 ± 0.46
Ours (ProbPrompt) 80.53 ± 0.37 83.90 ± 0.23 80.23 ± 0.31 84.21 ± 0.26 79.01 ± 0.44 83.58 ± 0.59

Table 2: CIFAR-100 与 CUB-200 上的 10-任务持续学习结果对比:

方法 CIFAR-100 FAA (%) CIFAR-100 CAA (%) CUB-200 FAA (%) CUB-200 CAA (%)
Joint-Training(上限) 91.38 - 88.41 -
DualPrompt 79.81 ± 1.19 88.48 ± 1.32 65.01 ± 1.08 77.56 ± 0.84
CODA-P 81.03 ± 0.78 84.26 ± 0.84 73.44 ± 0.62 81.55 ± 0.70
VQ-Prompt 88.73 ± 0.27 92.84 ± 0.73 86.72 ± 0.94 90.33 ± 1.03
APT 88.85 ± 0.63 92.84 ± 0.59 78.50 ± 0.94 -
Ours (ProbPrompt) 89.38 ± 0.22 93.34 ± 0.51 87.52 ± 0.41 90.93 ± 0.70

消融实验

Table 4 & 6 & 8: ImageNet-R 10-任务设置下各核心模块与超参数消融分析:

变体配置 / 采样数 \(N_s\) 机制说明 FAA (%) CAA (%)
纯确定性基线(\(\Sigma = \mathbf{I}\), 仅均值) 移除采样与方差学习,退化为确定性向量 77.49 ± 0.27 80.88 ± 0.43
独立单分布采样(无混合 \(\mathcal{N}_{\text{GM}}\) 独立采样各分布后由相关度得分加权 78.73 ± 0.51 82.46 ± 0.47
混合分布采样(无 \(\mathcal{L}_{\text{DR}}\) 正则化) 混合高斯分布 + 查询重加权聚合 79.98 ± 0.49 83.63 ± 0.39
简单平均聚合(\(w_k = 1/N_s\)\(\mathcal{L}_{\text{DR}}\) 忽略查询余弦权重,候选直接平权 79.06 ± 0.58 83.21 ± 0.51
无采样推理(直接使用混合均值 \(\mu_{\text{GM}}\) 训练或测试阶段剥离随机扰动 79.18 ± 0.18 82.54 ± 0.20
采样数 \(N_s = 1\) 仅抽样单个提示向量 无法收敛 (Fail) -
采样数 \(N_s = 5\) 样本量严重不足,随机方差过大 33.55 ± 13.61 36.57 ± 11.24
采样数 \(N_s = 10\) 训练波动仍然较大 61.04 ± 7.19 62.16 ± 8.05
完整模型 (\(N_s = 30\), 含 \(\mathcal{L}_{\text{DR}}\)) 混合分布 + 加权聚合 + 时序正则化 80.23 ± 0.31 84.21 ± 0.26

Table 7: 推理与训练计算开销对比(ImageNet-R 10-任务,ViT-B/16): - L2P: 显存 342.76 MB,训练 5.64 min/epoch,推理速度 74.54 FPS,FAA 69.29% - CODA-P: 显存 352.59 MB,训练 6.09 min/epoch,推理速度 74.38 FPS,FAA 75.45% - VQ-Prompt: 显存 343.87 MB,训练 6.30 min/epoch,推理速度 74.15 FPS,FAA 78.71% - Ours: 显存 347.85 MB,训练 6.33 min/epoch,推理速度 74.03 FPS,FAA 80.23%

关键发现

  • 打破提示坍塌的决定性因素:从纯确定性基线(77.49%)引入高斯分布建模(78.73%)乃至自适应混合分布(79.98%),准确率发生显著跃升。余弦相似度分析(Fig. 1(b))证实,ProbPrompt 的两两提示相似度降至全场最低(约 0.80 左右,大幅低于 baseline 的 0.88–0.96),证明表征坍塌已被实质性化解。
  • 采样数量的相变拐点\(N_s\) 从 1 增加到 20–30 呈现出阶跃式性能跃迁,从无法收敛直接飞跃至 80.23%;继续增加至 40(79.96%)性能基本饱和。这证明蒙特卡洛采样的主要作用在于通过适度随机性拓宽表征探索空间,并在样本均值聚合中稳定梯度估计。
  • 高效率与轻量化优势:与确定性方法相比,单前向批次推理开销仅增加不到 0.1 FPS,显存增加不足 4 MB,完全规避了引入大型生成网络或多前向传递的重度开销。

亮点与洞察

  • 提示表征空间的概率化升维:不同于以往在特征输出端加入方差噪声的折中做法,该方法直接在提示参数本身建立高斯分布并利用马氏距离建立相似度匹配,为解决参数高效微调(PEFT)中的特征同质化和坍塌问题提供了极具普适性的优雅方案。
  • 前向注入前完成降维融合:通过将 \(N_s\) 个采样的随机提示直接在输入层通过查询相似度压缩为单一聚合向量,兼得随机探索的表示多样性与单次前向的极致计算效率。
  • 分布平滑化解决可塑性-稳固性困境:利用跨训练步的 KL 散度约束作为记忆锚,在保证概率提示自由探索新特征的同时,严格限制了先验分布的剧烈畸变,成功阻止了灾难性遗忘。

局限与展望

  • 基准局限与骨干依赖:当前验证主要集中于预训练 ViT 骨干网络以及标准的闭集类别增量分类场景,尚未拓展至纯开放世界增量学习(Open-World Continual Learning)或多模态大语言模型(MLLM)的长序列微调任务。
  • 高斯对角假设的简化:为了控制计算开销,协方差矩阵采用了对角独立假设,忽略了提示各维度之间的潜在协方差结构;未来若结合低秩协方差分解可能进一步提升表达上限。
  • 采样依赖经验超参:采样数量 \(N_s\) 目前依靠网格搜索确定,在极度资源受限的边缘端无法自适应动态退火。

相关工作与启发

  • vs L2P / DualPrompt: L2P 使用最近邻选择检索不可微提示向量,DualPrompt 将提示硬拆为 G-Prompt 与 E-Prompt;两者均采用确定性离散向量,极易出现提示趋同。本文采用连续概率分布与连续马氏加权,全流程端到端可微,表达容量大幅超越确定性提示池。
  • vs CODA-P: CODA-P 引入注意力权重矩阵实现确定性组件的线性加权;本文则在参数空间定义分布,不仅考虑了均值还显式建模不确定性方差,利用随机扰动天然具备更强的抗过拟合与特征抗坍塌性能。
  • vs VQ-Prompt: VQ-Prompt 通过向量量化机制学习离散码本;而 ProbPrompt 则从连续高斯混合模型切入,无须维护复杂的码本更新与梯度直通估计(STE),结构更加纯粹且更易训练。

评分

  • 新颖性: ⭐⭐⭐⭐☆ [首次将提示参数概率化建模引入类增量持续学习并揭示提示坍塌现象,构思清晰新颖]
  • 实验充分度: ⭐⭐⭐⭐⭐ [覆盖 ImageNet-R/CIFAR/CUB 等多个任务切分,包含细致的消融、开销分析及自监督骨干泛化验证]
  • 写作质量: ⭐⭐⭐⭐⭐ [逻辑严密流畅,问题定义精准,数学推导自洽且图表翔实规范]
  • 价值: ⭐⭐⭐⭐☆ [计算开销极低且完全即插即用,为 ViT 持续学习和 PEFT 提示优化提供了强有力的新基线]