A Kernel-Based View of Language Model Fine-Tuning
Sadhika Malladi, Alexander Wettig, Dingli Yu, Danqi Chen, Sanjeev Arora
ICML 2023
cs.LG, cs.CL
2022-10-12
把预训练模型微调写成核回归问题,在 14 个 NLP 任务上验证 prompt 微调多数呈现核动力学,并从核视角解释 LoRA 为何有效。
预训练加微调这件事,从业者天天在做,却一直缺一个能解释它的理论。最朴素的反常是统计上的:一个上亿参数的语言模型,只用几十条标注样本去微调,按经典直觉早该过拟合得一塌糊涂,可它偏偏工作得很好。还有一连串顺手拈来却没人讲清的现象:同样是微调,加一句「It was [great/terrible]」的 prompt 与不加,效果差出一截;把参数更新限制在一个低秩子空间里(LoRA 这类做法),居然能追平全参数微调。
这些现象都指向同一个问题:微调到底在参数空间里走了多远?是大幅改写模型,还是只做了一点小修正?普林斯顿的这支工作(ICML 2023)给出的答案是,在相当多的情况下,微调更接近一次核回归(kernel regression),模型本身几乎没动。
神经切核(NTK)原本是研究无限宽、随机初始化网络梯度下降动力学的工具。核心直觉是:网络足够宽时,训练过程中每个参数的梯度几乎不变,于是整个训练可以被一个固定的核矩阵 K(ξ,ξ′)=⟨∇f(ξ),∇f(ξ′)⟩ 刻画,分类问题退化成对这个核做回归。可这套理论有两个前提不适用于微调:它要求随机初始化,而微调的起点是预训练好的、非随机的权重;它描述的是 SGD,而真实微调几乎都用 Adam。
作者分两步补上这个缺口。
第一步,给 Adam 造一个核。Adam 的全程动态很难写出核形式,因为每一步更新依赖整段梯度历史。但早期训练里,Adam 的矩估计只在邻域里滑动,更新退化成对梯度逐坐标取符号,即 SignGD。于是作者提出 SignGD 核:把标准 NTK 里的梯度换成梯度的符号,K^SignGD(ξ,ξ′)=⟨sign(∇f(ξ)),sign(∇f(ξ′))⟩,并证明它是 SignGD 的正确核类似物。这一步的意义在于,微调几乎都发生在这个早期阶段,所以 SignGD 近似够用。
第二步,解释为什么「非随机初始化」不破坏核行为。关键概念是「自然任务」(natural task):如果加合适的 prompt 后,下游任务在语义上就是预训练任务(掩码预测)的一个子类,那么当网络趋于无限宽时,预训练模型本身就已经接近能解这个任务,微调只需要做一个小修正,核行为成立。作者用 Tensor Programs 框架把这件事形式化(定理 5.5)。这同时给出一个干净的判据,prompt 之所以有用,是因为它把下游任务重新表述成「填空」,让它够得上「自然」。
实验在 14 个 NLP 任务上做(情感、话题分类、自然语言推理、复述检测),用 RoBERTa-base,每个任务抽 5 组 k-shot 数据,k 取 16 和 64。
主要结论是:eNTK(经验神经切核,用预训练权重直接算出来的核)能像微调一样解其中 12 个任务;8 个任务在微调过程中真正表现出完整的核行为(既满足线性化,也满足梯度近似不变)。带 prompt 时,核类似物在 10 个任务上与全微调的差距压到 10 个百分点以内。
prompt 的作用是决定性的。不带 prompt 的标准微调里,核与真实微调最多差 16 个百分点;带 prompt 后这个差缩到约 3 个百分点。同一个核,有没有 prompt 判若两物。
| 设置 | 核与微调的准确率差距 |
| 标准 FT(无 prompt),k=16/64 | 最高约 16 个百分点 |
| prompt-based FT | 约 3 个百分点 |
| prompt-based,核类似物 vs 全微调(10/14 任务) | <10 个百分点 |
| Adam-FT 与 SGD-FT(带 prompt) | <4 个百分点 |
少数任务始终不出现核行为:TREC、MNLI、SNLI、QNLI、MPQA。作者归因到 prompt 设计,比如 MNLI 把中性标签写成「Maybe」塞进句子里会造成不合语法的句子,任务就没法被表述成自然的填空。这从反面印证了「自然任务」假说。
第 7 节顺手解释了 LoRA。LoRA 把权重更新写成两个瘦矩阵的乘积 W+BA,本质是把全梯度投影到一个低秩随机子空间。由 Johnson-Lindenstrauss 引理,随机投影保持内积,于是核矩阵的每个元素都被保持,SGD-LoRA 在秩 k 足够大时拥有和全参数 SGD 完全一样的核与动力学。这是目前对「低秩微调为什么不掉点」最干净的理论解释之一。
这篇没有给出更强的微调方法,它给的是一个理解框架。从业者能拿走三件具体的东西。其一,prompt 不是玄学,它的作用是把任务对齐成预训练见过的形态,这点可以指导 prompt 设计。其二,LoRA 有效有了理论担保,只要秩够大,它不改变优化动力学,这给「秩该选多大」提供了一个角度。其三,eNTK 本身是一个不靠梯度更新就能用预训练模型解下游任务的方法,在标注极少、梯度噪声大的场景有时甚至比真实微调更稳。
这些结论是在 RoBERTa 级别的掩码模型上得到的。今天数百亿参数的自回归模型是否同样落在核行为区间,论文没有覆盖。
作者自己划了几条明确的边界。实验只覆盖少样本分类、单一掩码语言模型、特定的一组 prompt;要把 k 调大或换更大模型,eNTK 的计算开销会变成瓶颈。理论上 SignGD 近似只刻画 Adam 的早期阶段,长训练是否还成立并不清楚,同期工作(Littwin & Yang 2023)暗示「Adam 归约到 SignGD」是能否看到核行为的关键,意味着这套框架的适用窗口比看上去窄。
核行为的判定阈值(线性化改进过半、梯度距离小于 2.0)是作者手选的,论文也坦承形式定义不规定数值阈值。8/14 这个数字对这些阈值敏感,换一套阈值,结论的强弱会跟着动。好在作者把原始数据摊在表里,读者可以自己核。