跳转至

Estimating and Orthogonalizing Unknown Pre-training Gradients for Continual Fine-tuning of Large Language Models

会议: NeurIPS2026
arXiv: 2609.30935
代码: https://github.com/wangbing1416/EoupCT
领域: 优化/理论;大语言模型持续微调
关键词: 持续学习、知识保持、可微伪数据、软提示、梯度正交化

一句话总结

EoupCT 用冻结预训练模型和可学习软提示构造容易被新任务破坏的可微知识代理,再联合蒸馏与梯度投影保护历史任务和通用能力;在六个模型的 SuperNI/MMLU 子集实验中改善遗忘,但其“纯一阶”和“绝对零遗忘”表述需要限定。

研究背景与动机

持续微调面对两种不同的遗忘:刚学完后续任务,早期下游任务的表现可能下降;即使早期任务保持得不错,预训练阶段获得的通用知识也可能逐渐流失。历史任务的训练数据或梯度通常可记录,但现成模型的原始预训练语料与训练轨迹不可得。因而,仅将新任务更新限制在历史下游梯度的正交空间,不能自动保护没有进入这个空间的预训练能力。

已有方法各有覆盖盲区。OLoRA 约束不同任务的低秩更新子空间,CLoRA 通过子空间正则限制输出扰动,LoRAMoE 和 GainLoRA 用多分支及路由减少任务间干扰。它们可以降低持续适应中的冲突,却不等于知道当前任务会破坏哪些通用知识。随机生成一些文本进行回放也不一定抓住这种薄弱处:容易生成的知识可能已经稳定,而真正需要保护的方向可能很少被采样到。

本文把预训练模型本身作为知识代理的来源,并让软提示随当前任务变化,主动寻找学生模型在一次新任务更新后更难维持的表征。这里的“估计预训练梯度”是学生对教师代理蒸馏损失的梯度,不是恢复训练时记录的真实梯度,也不需要取得原始预训练数据。核心 idea:用任务相关的可微伪数据找出易遗忘的知识方向,再将这些方向与历史任务梯度共同纳入持续微调的更新约束。

方法详解

整体框架

输入是冻结的预训练教师、从其初始化的可训练学生,以及顺序到来的下游任务;实际实验只训练 LoRA 参数。每个任务配置一组软提示,通过教师自回归地产生连续伪嵌入序列。教师与学生在同一代理序列上的表征差异提供知识保持信号,而新任务的监督损失提供学习信号。

训练依次经过“可微知识代理”“虚拟步脆弱性搜索”“历史零空间投影”“冲突修正与联合更新”。前两步调整提示并计算保护梯度,后两步调整学生更新方向。冻结教师意味着不更新其权重,并不意味着生成链可以整体停止求导:更新提示仍需要梯度穿过教师对输入的计算。

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    T["冻结预训练教师"] --> A["可微知识代理"]
    P["可学习软提示"] --> A
    A --> B["虚拟步脆弱性搜索"]
    N["新任务监督梯度"] --> B
    B -.->|只优化提示| P
    B -->|真实学生权重上的保护梯度| C["历史零空间投影"]
    N --> C
    H["历史下游梯度"] --> C
    C --> D["冲突修正与联合更新"]
    D --> S["更新学生 LoRA"]

图中回边是训练时的提示优化,底部才是学生参数更新;部署学生不需要重新执行这条伪数据生成链。历史保护梯度用于提示多样性惩罚,历史下游梯度用于学生的零空间投影,两者不能混为同一个记忆库。

关键设计

1. 可微知识代理:把不可求导的文本采样替换为连续伪嵌入

直接生成离散词元会切断提示到生成序列的梯度。作者为每个任务设置长度为 \(L\) 的软提示,并让冻结教师生成长度为 \(M\) 的序列。在每一步,将词表 logits 加上 Gumbel 噪声,再用温度控制的 softmax 得到权重;下一步输入不是某个硬词元,而是词嵌入的加权和。

\[ \pi_{m,i}=\frac{\exp((z_{m,i}+\varepsilon_{m,i})/\tau)}{\sum_j\exp((z_{m,j}+\varepsilon_{m,j})/\tau)},\qquad \mathbf e_m=\sum_i\pi_{m,i}(\mathbf W_{\mathrm{emb}})_i. \]

这使提示能够通过整个自回归生成链影响最终的保护损失。低温度会使权重更接近 one-hot,但存在 Gumbel 噪声时,它对应带噪声的随机类别选择,不是确定性贪心解码。附录 A.1 将极限解释为贪心词元的部分需要谨慎,不能据此宣称恢复了真实预训练文本分布。

教师和学生在这些伪嵌入上的潜在表征采用均方误差蒸馏;所谓保护梯度是该损失对学生可训练参数的导数。它测量当前学生偏离教师的方向,与原始预训练目标、语料梯度及训练历史均不是同一对象。

\[ \mathcal L_{\mathrm{pre}}(\boldsymbol\theta;\mathbf S(\mathbf P_t))=\ell_{\mathrm{MSE}}\big(\mathcal F_{\boldsymbol\theta^0}(\mathbf S(\mathbf P_t)),\mathcal F_{\boldsymbol\theta}(\mathbf S(\mathbf P_t))\big),\qquad \mathbf g_{\mathrm{pre}}(\mathbf P_t)=\nabla_{\boldsymbol\theta}\mathcal L_{\mathrm{pre}}. \]

附录 D.5 实际只在每一步概率最高的 50 个候选词元上进行松弛,而非始终在整个词表上计算。它降低存储需求,却也限制了代理能覆盖的词元;候选集合切换处不是全局光滑的生成映射。因此,应将代理理解为受提示、长度和候选截断约束的局部知识探针。

2. 虚拟步脆弱性搜索:让提示针对新任务造成的表征漂移

仅最大化当前学生与教师的差异,会找到学生已经遗忘的内容,却未必找到当前新任务即将破坏的内容。作者先计算新任务梯度,暂时把学生沿该梯度移动一步,再在这个虚拟学生上优化提示。真实学生此时尚未接受该步更新;提示搜索结束后,保护梯度仍在真实学生权重处重算。

\[ \boldsymbol\theta_{\mathrm{virt}}=\boldsymbol\theta-\alpha\mathbf g_{\mathrm{new}},\qquad \mathbf P_t^*=\arg\max_{\mathbf P_t}\left[\mathcal L_{\mathrm{pre}}(\boldsymbol\theta_{\mathrm{virt}};\mathbf S(\mathbf P_t))-\lambda\sum_{i<t}\cos^2\big(\mathbf g_{\mathrm{pre}}(\mathbf P_t),\mathbf g_{\mathrm{pre}}(\mathbf P_i)\big)\right]. \]

第二项惩罚当前保护梯度与历史保护梯度的平方余弦相似度,目的是不要每个任务都重复寻找同一块知识。它鼓励方向多样性,但有限强度的软惩罚不保证严格正交,也不证明语义主题或整个预训练分布得到充分覆盖。

作者用一阶泰勒展开解释虚拟步为何能暴露冲突。下面显式保留提示依赖,属于对原文式 (8) 的记号展开,而非新增算法。

\[ \mathcal L_{\mathrm{pre}}(\boldsymbol\theta-\alpha\mathbf g_{\mathrm{new}};\mathbf S(\mathbf P_t))\approx\mathcal L_{\mathrm{pre}}(\boldsymbol\theta;\mathbf S(\mathbf P_t))-\alpha\langle\mathbf g_{\mathrm{new}},\mathbf g_{\mathrm{pre}}(\mathbf P_t)\rangle. \]

若虚拟权重视为固定值,第一项虚拟损失对提示确实可用普通反向传播优化,无须对新任务梯度再求导。但展开中的基线损失也依赖提示,所以最大化虚拟损失并不无条件、精确等价于最大化梯度冲突;它还可能偏好原本蒸馏误差就大的代理。

更重要的是,完整目标包含对保护梯度的余弦惩罚。保护梯度本身依赖提示,对该项求导仍涉及提示与学生参数的混合二阶导数。固定虚拟权重只解释了虚拟损失项的求导,不能自动把完整目标变成纯一阶。原文未交代此项是否使用停止梯度、近似或替代目标,笔记不假定代码已经消除了这个问题。

3. 历史零空间投影:先守住已记录的下游任务方向

当前任务梯度和当前保护梯度都可能影响早期下游任务。作者把历史下游梯度按列拼成矩阵,将两个当前梯度共同投影到其正交补空间。先投影两者而不是只投影新任务梯度,避免后续加入蒸馏更新时重新引入历史任务冲突。

\[ \boldsymbol\Pi_{\mathrm{new}}=\mathbf I-\mathbf M_{\mathrm{new}}(\mathbf M_{\mathrm{new}}^\top\mathbf M_{\mathrm{new}})^{-1}\mathbf M_{\mathrm{new}}^\top,\qquad \tilde{\mathbf g}_{\mathrm{new}}=\boldsymbol\Pi_{\mathrm{new}}\mathbf g_{\mathrm{new}},\quad \tilde{\mathbf g}_{\mathrm{pre}}=\boldsymbol\Pi_{\mathrm{new}}\mathbf g_{\mathrm{pre}}^*. \]

原文式 (9) 的普通逆要求历史梯度列线性独立。空历史记忆时应理解为没有历史约束;历史方向重复或退化时该逆可能不存在,需要合适的数值处理,但论文未说明具体实现,不能把伪逆或正则化写成作者已经采用的方案。

算法 1 在每个任务结束后向两种记忆库各加入一个梯度,未明确其是否为全任务平均、最后一个批次或另行估计。保护的是被记录的方向,而不是所有历史样本在所有后续权重处的梯度。随着任务增多,记忆存储和被排除的子空间都会增长,可能减少学习新任务的自由度。

4. 冲突修正与联合更新:只删去新任务中破坏当前代理的分量

经过历史投影后,作者检查新任务梯度与保护梯度的内积。非负意味着沿梯度下降方向学习新任务不会在一阶上增加代理损失;负内积则意味着冲突,需要从新任务梯度中移除沿保护梯度的冲突分量。最后把修正后的新任务梯度与保护梯度相加。

\[ \tilde{\mathbf g}_{\mathrm{new}}^*=\begin{cases}\tilde{\mathbf g}_{\mathrm{new}}-\dfrac{\langle\tilde{\mathbf g}_{\mathrm{new}},\tilde{\mathbf g}_{\mathrm{pre}}\rangle}{\|\tilde{\mathbf g}_{\mathrm{pre}}\|^2}\tilde{\mathbf g}_{\mathrm{pre}},&\langle\tilde{\mathbf g}_{\mathrm{new}},\tilde{\mathbf g}_{\mathrm{pre}}\rangle<0,\\\tilde{\mathbf g}_{\mathrm{new}},&\text{otherwise},\end{cases}\qquad \boldsymbol\theta^+=\boldsymbol\theta-\eta(\tilde{\mathbf g}_{\mathrm{new}}^*+\tilde{\mathbf g}_{\mathrm{pre}}). \]

附录 A.2 证明的是冲突情况下,欧氏距离最近的正交梯度修正,不是整个非凸训练问题的全局 Pareto 最优解。联合更新包含降低代理蒸馏损失的分量,所以最终更新不必与当前保护梯度严格正交;真正需要的是代理损失的一阶变化不为正,以及对已记录历史方向保持正交。

这些性质依赖当前梯度、精确投影和对应的实际更新方向。它们是局部的一阶保证,不能推出非线性模型有限步更新后输出不变或绝对零遗忘。附录 D.5 又采用 AdamW;若投影之后再施加逐坐标自适应缩放、动量或权重衰减,实际参数增量未必仍处在同一零空间,原文没有解释如何让优化器与投影保证一致。

损失函数 / 训练策略

新任务使用目标回答的自回归负对数似然,学生同时承担教师表征的蒸馏目标。提示做最大化,学生做最小化,两者不能理解为同一损失一次反向传播更新所有参数。每次迭代按新任务梯度、虚拟提示搜索、真实保护梯度、两级投影、学生更新的顺序执行;任务结束后再扩充记忆。

如果教师和学生完全相同,均方误差及其学生梯度可能为零;投影也可能将非零保护梯度压成零。此时平方余弦的归一化可能未定义,冲突投影的分母也必须按非零条件解释。论文没有完整给出这些退化情形的初始化与数值防护,不能从公式推断其实现细节。

实验采用 LoRA rank 4、AdamW 学习率 \(2\times10^{-4}\)、梯度检查点、训练输入上限 1024 词元和回答生成上限 50 词元。MMLU 使用最多 5-shot、上下文上限 2048;超出长度则减少示例,批大小为 1,并以候选答案的末位置 logits 选择答案。

提示长度与代理长度的敏感性范围是 \(\{4,8,16,32\}\),附录举例使用代理长度 16 或 32。原文没有完整披露所有模型的提示优化内循环、温度日程、虚拟步长和多样性系数,不能把敏感性曲线当作统一默认配置。

实验关键数据

主实验

持续学习流为 SuperNI 的 15 个任务,涵盖问答、信息抽取、情感分析、摘要与对话,使用三种随机任务顺序。通用能力评测是 MMLU 的 9 个学科,分为 STEM、Humanity 与 Other,而非已经证实覆盖完整 MMLU。表 1 的 MMLU 数字是三个任务顺序的平均,SuperNI 则逐顺序报告。

SuperNI 用 ROUGE-L 衡量生成质量。其 Fgt. 是各任务“刚学完该任务的分数减去全部任务学完后的分数”的平均;MMLU Fgt. 则为原始模型准确率减最终准确率。两者都是分数差,不是相对百分比,负遗忘值表示最终得分高于对应基线。

\[ \mathrm{FR}_{\mathrm{SuperNI}}=\frac1T\sum_{t=1}^{T}(R_{t,t}-R_{T,t}),\qquad \mathrm{FR}_{\mathrm{MMLU}}=A_0-A_T. \]

下表摘录原文表 1 的 Order 1 与 STEM 列,保留不同模型上的代表性对照;它不是所有顺序和 MMLU 学科的汇总表。

模型 方法 SuperNI ROUGE-L ↑ SuperNI Fgt. ↓ MMLU STEM Acc. ↑ MMLU STEM Fgt. ↓
Qwen3-4B LoRA 45.1 13.7 56.4 4.5
Qwen3-4B OLoRA 48.9 9.0 58.1 2.9
Qwen3-4B EoupCT 50.7 6.0 60.2 0.7
Llama3-8B OLoRA 44.6 12.8 43.6 2.4
Llama3-8B EoupCT 51.4 6.5 45.0 1.0
Gemma2-9B OLoRA 50.0 8.8 42.9 10.3
Gemma2-9B EoupCT 53.4 4.8 52.3 0.9

例如 Qwen3-4B 相对 OLoRA 的任务分数提高 1.8 分,任务遗忘差降低 3.0 分;Gemma2-9B 的 STEM 准确率提高 9.4 个百分点。这里的结果支持在该协议下改善知识保持,不证明未知预训练梯度得到了准确恢复。

消融实验

下表为原文表 2 的 Qwen3-4B 列。C2 是脆弱性搜索,P&PO 是保护梯度之间的多样性约束,N&PO 是新任务与保护梯度的冲突约束,N&NO 是历史下游梯度约束。表 2 将 MMLU 列简写为 Acc./Fgt.,完整模型数值与表 1 的 STEM 列一致;因此保留原表标签,不把它解释为全部 9 个学科的平均。

配置 SuperNI ROUGE-L ↑ SuperNI Fgt. ↓ MMLU Acc. ↑ MMLU Fgt. ↓
EoupCT 50.7 6.0 60.2 0.7
w/o C2 47.6 10.3 59.0 2.0
w/o P&PO 47.3 9.2 57.8 3.1
w/o N&PO 47.4 8.7 57.9 3.0
w/o N&NO 46.9 10.3 56.9 4.1

移除脆弱性搜索后,任务分数下降 3.1 分;移除历史下游保护后,任务分数下降 3.8 分,MMLU 遗忘差从 0.7 增至 4.1。它说明定向寻找代理与保护历史任务都有作用,但不能只凭消融得出完整目标已经严格满足理论约束。

附录表 3 报告同一 8-GPU 设置下的运行分钟数。保留分钟而非重算作者取整后的倍率,以免将不同列的相对成本混淆。

模型 LoRA OLoRA CLoRA EoupCT
Qwen3-4B 62 204 155 129
Llama3-3B 35 151 114 65
Gemma2-9B 76 271 210 178
作者报告的平均 60 203 157 122

EoupCT 在所报配置中比 OLoRA 和 CLoRA 快,但明显慢于普通 LoRA;平均 122 分钟对 60 分钟,不能概括为“几乎没有训练开销”。该表没有充分披露 GPU 型号、峰值显存及统一预算细节,也不能直接外推单卡或更长代理的成本。

关键发现

  • 三种任务顺序和三个模型家族的结果共同支持方法的稳健性,但不能替代固定顺序下多次独立训练的误差条或显著性检验。
  • 提示和代理过短时表现较弱,增大后改善,达到 32 时可能略降;缓存只有曲线说明,不补写读不到的精确曲线数值。
  • 表 1 的 Llama3-3B STEM Fgt. 为 -0.8,说明该子集可能发生正迁移,但不能据此推断所有通用知识都得到增强。
  • “移除任何组件都使所有遗忘指标变差”不是逐格成立:表 2 的 Gemma2-2B 在 w/o C2 下 SuperNI Fgt. 为 10.1,完整模型为 10.3。任务分数整体下降与每个遗忘数值单调恶化必须区分。

亮点与洞察

  • 保护对象由任务决定。 伪数据不只是便宜的历史回放材料,而是用于发现当前更新最容易破坏的教师—学生差异;这使知识保持具有任务针对性。
  • 两种历史记忆承担不同职责。 历史保护梯度用于扩大知识探针的方向覆盖,历史下游梯度直接限制学生更新。把它们分开,能看清“代理多样性”与“旧任务稳定性”并非同一目标。
  • 先投影、后处理冲突的顺序有意义。 两个梯度先进入同一个历史零空间,再修正彼此冲突,避免局部知识保护重新破坏已有下游约束。

局限与展望

  • 知识代理不等于预训练分布。 短提示、短序列和 top-50 截断限制覆盖,代理可能处在软嵌入空间而非自然文本空间;教师的生成偏好也可能漏掉不易生成的知识。可以进一步比较自然文本回放与软嵌入保护,并扩大评测范围。
  • 理论边界强于实际证据。 附录 A.3 的全局界依赖提示覆盖真实分布、找到最坏提示,并在所需权重处控制最坏损失。虚拟权重处的最坏提示不自动是更新后权重处的最坏提示;Remark 3.1 的局部代表性解释也没有补足全局覆盖证明。
  • 一阶效率仍有未解条件。 完整梯度余弦惩罚的提示导数涉及混合二阶项;附录 B.1 把 Hessian-vector product 一概描述为平方级显存爆炸并不成立,自动微分可不显式形成 Hessian。附录 B.2 又把平方复杂度称为“指数”,这也是不准确的复杂度术语。
  • 退化与实际更新需要明确。 零保护梯度、历史 Gram 矩阵不可逆、AdamW 后的实际增量以及过时历史方向,都关系到投影约束是否真实成立。未来需要披露数值处理并测量实际更新的约束残差,而不是只引用形式上的正交性。
  • 实验结论仍有限。 只有短任务流、有限模型规模和 MMLU 子集,未报告充分的显著性统计。正文称组件删除一致恶化,而表 2 存在反例;附录 D.2 同时使用子集与“full MMLU”措辞,范围需要作者澄清。
  • 复现信息不够完整。 作者列出的“Llama3-3B”未给出精确检查点,不能擅自改称另一个模型版本;提示优化细节和显存配置也未充分提供。附录 E 的动态矩阵与表头汇总不宜未经对账便当作同一平均口径。

相关工作与启发

  • vs OLoRA / CLoRA:它们主要从低秩权重或子空间限制扰动;EoupCT 另外构造任务相关的教师代理并调整梯度方向,代价是生成、提示搜索与保护梯度计算。
  • vs LoRAMoE / GainLoRA:多专家及门控通过参数分工降低冲突;EoupCT 强调单个学生更新时的知识约束。运行时间优势只适用于表 3 的具体配置,不代表所有路由方案更昂贵。
  • vs SSR / LAMOL:伪回放通常合成历史任务样本再联合训练;这里生成的是可微嵌入代理,并根据虚拟更新搜索脆弱方向。可借鉴的是任务定向的保留探针,而不是“已经恢复原始训练数据”的结论。

评分

  • 新颖性: 4/5 — 将可微知识代理、虚拟步搜索与持续梯度投影串联,问题定位明确。
  • 实验充分度: 3/5 — 多模型、多顺序及消融较完整,但通用能力覆盖、统计和复现细节有限。
  • 写作质量: 2/5 — 方法流程可理解,纯一阶、全局保护和绝对零遗忘的论述明显过强。
  • 价值: 4/5 — 为无需原始预训练数据的持续微调提供了可研究的知识保持方向,仍需实现与理论校准。