ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

GEM持续学习:梯度投影与二次规划防遗忘

GEM持续学习:梯度投影与二次规划防遗忘 1. 为什么持续学习绕不开 GEM 这道坎做 Continual Learning 的人早晚都会撞上 Gradient Episodic MemoryGEM这个名字。它是 2017 年 NeurIPS 上的工作作者是 David Lopez-Paz 和 MarcAurelio Ranzato放到今天看依然是被引用最多、被拿来当 baseline 最多的经典方法之一。我最初接触它是因为手上的模型在做增量训练时换了新数据后老任务的准确率掉得像跳水一夜回到解放前。GEM 给出的思路很朴素既然遗忘是因为参数被新任务带跑偏了那我就给梯度的更新方向加个约束别让新任务的步子踩到老任务的地盘上。适合读这篇的你应该已经写过至少一个分类网络懂反向传播也大概知道什么是梯度剩下的我尽量说清楚。有个小插曲得先说一下。如果你在搜索引擎里敲 GEM很可能搜到一堆半导体设备通信的 SECS/GEM 标准那玩意儿和咱们聊的 Gradient Episodic Memory 完全是两码事缩写撞车了而已。本文里的 GEM 一律指持续学习里的梯度情景记忆别被带偏。先说清楚 GEM 想解决的核心矛盾。常规的神经网络训练假设数据是独立同分布的你可以把全部任务的数据混在一起打乱再训练。但在持续学习场景里任务是一个接一个来的比如先学十个类别的图像分类再学另外十个类别而且旧数据因为存储或隐私原因拿不到了。这时候如果你只用新任务的数据做梯度下降网络参数会朝着新任务的最优解狂奔结果就是在旧任务上表现崩盘。这个现象叫灾难性遗忘Catastrophic Forgetting是持续学习的头号敌人。GEM 的定位很明确它属于基于记忆回放这一大类方法通过保留每个旧任务的极少量样本在学习新任务时用这些样本约束梯度方向从而在允许正向知识迁移的同时压制负向的遗忘。为什么这个约束要设计成梯度内积不小于零而不是损失不上升这是理解 GEM 的关键。损失不上升是个数值条件很难直接写进优化里而梯度内积是个几何条件天然可以和梯度下降的框架结合还能化成一个标准的二次规划问题去解。这个转化是 GEM 最漂亮的地方后面我单独用一节拆开讲。2. GEM 的数学内核把遗忘问题变成一个二次规划2.1 从回放旧数据到约束旧梯度最直觉的做法是经验回放把旧任务的样本和新任务的样本混在一个 batch 里一起训练。这招确实有效但它有两个问题。第一你需要存足够多的旧样本否则回放的效果不稳定第二混着训练并不能保证旧任务的性能不下降你只是指望数据分布被拉平了而已没有硬约束。GEM 的改进点在于它不满足于混着训而是要显式地保证新任务更新完之后旧任务的损失不会变大。具体怎么保证作者用了一个一阶泰勒展开的近似。假设当前参数是 θ旧任务 k 的损失是 L_k(θ)如果参数更新一个很小的量 Δ那么 L_k(θΔ) ≈ L_k(θ) ⟨∇L_k(θ), Δ⟩。我们希望更新后旧任务损失不上升也就是 L_k(θΔ) ≤ L_k(θ)在步长足够小的前提下这就近似等价于 ⟨∇L_k(θ), Δ⟩ ≤ 0。又因为 Δ 正比于新任务的梯度 g方向上取反最终落到梯度层面就是一个干净的条件新任务的梯度 g 和旧任务的梯度 g_k 的内积要大于等于零即 ⟨g, g_k⟩ ≥ 0。这个内积大于等于零的几何含义是什么它意味着新梯度在旧梯度方向上不能有反向分量。如果两者夹角小于 90 度说明新任务的更新对旧任务的损失是中性或有益的正向迁移一旦夹角超过 90 度新梯度就有把旧任务损失抬高的倾向GEM 就要把这个分量砍掉。注意这里用砍掉是不准确的GEM 做的是投影不是裁剪投影后得到的梯度方向是满足所有旧任务约束的条件下离原始新梯度最近的那个方向。2.2 约束条件的几何图景把所有旧任务的梯度看成一组向量 g_1, g_2, ..., g_{t-1}每个约束 ⟨g, g_k⟩ ≥ 0 定义了一个半空间也就是和 g_k 夹角不超过 90 度的方向集合。所有半空间的交集是一个凸锥convex cone也被叫做可行域。新任务的原始梯度 g_new 如果本来就落在这个凸锥里说明它不会伤害任何一个旧任务直接用它就行如果它在锥外我们就要把它投影到锥上找到锥内距离它最近的那个梯度 g。这个投影到凸锥的说法很好用它能解释 GEM 为什么天然允许正向迁移。假如某个旧任务和新任务高度相关旧任务梯度方向和新任务梯度方向基本一致那么投影的约束几乎不生效新任务照样能大步往前走旧任务还能跟着受益这就是正向迁移的来源。反过来如果两个任务冲突严重投影会强行把新梯度掰到不伤害旧任务的方向上代价就是新任务学得慢一些。这种牺牲部分新任务学习速度换取旧任务不遗忘的权衡是 GEM 的固有特性你在调参时会反复体会到。2.3 把投影写成标准二次规划投影问题本身是一个有约束的优化我们可以把一个求最近点的问题写成minimize_g (1/2) ||g - g_new||^2 subject to G g 0其中 G 是一个 (t-1) 行 d 列的矩阵第 k 行就是旧任务 k 的梯度 g_kd 是参数维度。这个问题有约束、有二次目标正是标准的二次规划QP。但直接对 d 个变量求解太慢了d 可能是几百万。GEM 的巧妙之处在于转成对偶问题把变量维度从 d 降到任务数 (t-1)这样即使模型很大QP 的规模也只是任务数量级。对偶的推导不复杂我这里把结论说清楚。设拉格朗日乘子 v ≥ 0对原始问题关于 g 求导置零可以得到 g g_new G^T v。代回目标函数原始的最小化问题等价于求解下面这个对偶问题minimize_v (1/2) v^T (G G^T) v (G g_new)^T v subject to v 0解出 v 之后回代得到投影梯度 g g_new G^T v。注意 G G^T 是一个 (t-1) × (t-1) 的小矩阵Q 里面的每个元素就是两个旧任务梯度之间的内积 ⟨g_i, g_j⟩所以常常把它叫做梯度内积矩阵GEM 名字里的 Memory 和这个矩阵关系不大这里的 G 矩阵是 gradient 的意思别和整体方法名混淆。这个形式在实现时非常好写也是几乎所有开源复现的基础后面实操部分我会把它落成具体的代码。2.4 为什么是内积不小于零而不是损失不上升这里再强调一次因为很多新手会问。理论上如果我们能精确计算每个旧任务在新参数下的损失那直接拿损失当约束最稳妥。但损失是一个非线性函数写进 QP 里就变成非凸问题了没法解。梯度内积是一阶近似它只在步长足够小的情况下有效所以 GEM 必须用小学习率、小步长来保证近似的合理性。这也是为什么 GEM 训练时学习率通常要比普通训练小一截代价就是收敛慢。这个细节在论文里没有特别强调但我实际跑的时候感受很深学习率设大了内积约束的近似就崩了旧任务照样遗忘。3. 从零实现一个能跑的 GEM3.1 网络结构与记忆模块的骨架先把整体结构定下来。你需要三样东西一个分类网络比如一个小型 CNN 或 MLP视任务而定、一个情景记忆模块Episodic Memory缩写也是 EM注意别和 Expectation Maximization 混了、以及一个 QP 求解器。记忆模块的职责是按任务存储固定数量的样本每个任务存 m 个m 通常在 100 到 500 之间具体看任务难度和显存。记忆里的样本不是随机存的常见做法是任务训练结束后按类别均匀采样保证每个类都有代表避免某个类被淹没。网络本身用普通的分类网络就行GEM 不挑结构。我一般用两层卷积加两层全连接的小网络测试参数量控制在百万级方便调试 QP。记忆中样本的存储方式有两种存原始输入图像、文本特征或者存网络中间层的特征。存原始输入更通用能应付网络结构调整存特征省显存但换网络就得重算。我建议新手先存原始输入简单直接。样本的存储时机也值得说一句。GEM 论文里是在每个任务训练结束后从该任务的训练集里采样 m 个样本存入记忆。这个时机很讲究训练过程中存样本样本可能被当前模型过拟合过回放价值反而下降训练结束后存样本是经过完整训练后模型见过的用作约束更稳定。另外存储时最好打乱类别顺序再取防止取到的全是同一个类。3.2 记忆的采样与梯度计算细节每次训练新任务的一个 batch 时关键动作是从记忆里对每个旧任务采样一个 batch 的样本分别计算旧任务的梯度。这里有个很容易踩的坑就是旧任务梯度的计算必须用当前参数重新前向反向不能缓存。因为参数一直在变旧梯度也跟着变缓存下来的梯度是过期数据约束就失真了。这部分计算开销是 GEM 相比普通训练慢的根源我给你算一笔账。假设现在训到第 t 个任务记忆里有 t-1 个旧任务每个旧任务采样 b 个样本。如果新任务的 batch 大小是 B那么一次参数更新需要前向反向 B b×(t-1) 个样本。当 t10、b10 时额外开销是 90 个样本的反向传播大约是主 batch 的 9 倍假设 B 也是 10 左右。这就是为什么 GEM 在大规模任务序列下会很慢任务越多越慢是线性增长。想控制开销可以把 b 设得很小比如 5 到 10够用就行多了边际收益不明显。旧任务梯度的计算还有个维度问题。每个样本单独算梯度太贵GEM 用的是把一个任务的 b 个样本当一个 batch 算一个平均梯度。这样每个旧任务只得到一个梯度向量 g_k约束数就等于旧任务数。如果你非要把每个样本都当一个约束QP 规模会爆炸得不偿失。用任务级别的平均梯度是工程上的必要妥协效果也足够好。3.3 QP 求解代码怎么写下面给一个基于 cvxopt 的 GEM 核心实现片段。我把它写成函数方便你直接抄。注意 Q 矩阵和 p 向量的构造这是最容易出错的地方。import torch import cvxopt import numpy as np def project_gradient(g_new, grads_old): g_new: 新任务的梯度, 一维 tensor, shape (d,) grads_old: list of 旧任务梯度, 每个 shape (d,) 返回投影后的梯度 g, shape (d,) if len(grads_old) 0: return g_new # 拼成矩阵 G, shape (t-1, d) G torch.stack(grads_old, dim0) # Q G G^T, shape (t-1, t-1) Q torch.mm(G, G.t()) # p G g_new, shape (t-1,) p torch.mv(G, g_new) # 转成 cvxopt 需要的类型 t Q.size(0) P cvxopt.matrix(Q.double().cpu().numpy()) q cvxopt.matrix(p.double().cpu().numpy()) # 约束 v 0 G_cvx cvxopt.matrix(-np.eye(t)) h_cvx cvxopt.matrix(np.zeros(t)) # 关掉求解器输出 cvxopt.solvers.options[show_progress] False try: sol cvxopt.solvers.qp(P, q, G_cvx, h_cvx) v torch.tensor(np.array(sol[x]).flatten(), dtypeg_new.dtype, deviceg_new.device) except Exception: # 求解失败时退回原梯度, 后面会讲为什么 return g_new # g g_new G^T v g g_new torch.mv(G.t(), v) return g这段代码有几点要解释。第一Q 可能不是正定的旧梯度之间线性相关时cvxopt 的 qp 求解器本身对半正定也大体能处理但偶尔会失败所以外面套了 try。第二如果求解失败我选择退回到原始梯度 g_new让训练继续进行。这个策略叫fail-safe宁可这一次不做约束也别让训练直接崩后面排查章节会展开。第三所有张量运算要转到 double 再给 cvxoptfloat32 有时候会因为精度问题让求解器报错这点很多人卡过。还有个细节如果旧任务梯度数量很多比如超过 20 个QP 求解时间会明显上升。这时候可以考虑 GEM 的简化版 A-GEM它只用所有旧梯度的平均作为单一约束QP 退化成对一维向量的简单判断速度飞快。A-GEM 的取舍是约束弱一些但工程上友好太多我在任务数超过 20 时基本都会切到它。3.4 训练主循环的完整骨架把上面的模块拼起来训练循环长这样每个 batch 先算新任务的损失和梯度然后从记忆里采样算每个旧任务的梯度调用投影函数得到约束后的梯度最后把控制权交回优化器。注意这里有个微妙的地方优化器里已经有了一份新任务的梯度但那份梯度还没被投影过。我的做法是手动管理梯度不用 optimizer.step() 的默认流程而是把投影后的梯度写回参数的 .grad再调用 step。for task_id, (loader, mem) in enumerate(task_sequence): for x, y in loader: # 1. 新任务的前向反向 logits model(x) loss criterion(logits, y) optimizer.zero_grad() loss.backward() # 2. 收集新任务梯度 g_new [] for p in model.parameters(): g_new.append(p.grad.data.view(-1).clone()) g_new torch.cat(g_new) # 3. 算旧任务梯度 grads_old [] for k in range(task_id): bx, by mem.sample_batch(k) logits_k model(bx) loss_k criterion(logits_k, by) model.zero_grad() loss_k.backward() gk [] for p in model.parameters(): gk.append(p.grad.data.view(-1).clone()) grads_old.append(torch.cat(gk)) # 4. 投影 g_proj project_gradient(g_new, grads_old) # 5. 把投影后的梯度写回 offset 0 for p in model.parameters(): numel p.numel() p.grad.data.copy_( g_proj[offset:offsetnumel].view_as(p)) offset numel optimizer.step()有两处特别容易出 bug。一个是第 3 步里的model.zero_grad()你在算旧任务梯度前必须把上一步的梯度清掉否则不同任务的梯度会累加投影出来的约束完全错误。另一个是第 5 步的偏移量管理一定要保证参数顺序和拼梯度时的顺序完全一致我见过有人因为模型里加了 BatchNorm 层参数的遍历顺序和构建梯度时不一致结果投影后的梯度错位训练效果一塌糊涂。3.5 参数配置与调参经验给你一份我实测下来比较稳的起点配置你可以在此基础上微调。参数建议值说明记忆大小 m200 到 500每任务任务越难越大显存够就多存旧任务采样数 b10太小梯度噪声大太大速度慢学习率0.01 到 0.05比普通训练小保证一阶近似成立batch 大小10 到 30单任务 batch别设太大优化器SGD 或 AdamSGD 更贴合论文Adam 也能跑训练轮数每任务 10 到 30GEM 收敛慢要给足轮数学习率这件事值得单独说。GEM 的约束是基于一阶近似的步长一大泰勒展开就不成立弃掉的旧任务损失实际上还是会涨。所以我一般把学习率设成普通训练的 1/3 到 1/2。如果任务间冲突特别严重比如任务顺序是猫狗分类接着车辆分类再接着花卉分类这种差异大的序列学习率还得再降。你不要嫌它学得慢持续学习本来就是在稳定性和可塑性之间走钢丝GEM 偏保守这是个特点不是 bug。4. 排查实战GEM 跑起来会遇到的那些坑4.1 QP 求解失败与数值不稳定这是 GEM 复现里最高频的问题没有之一。表现是训练中途突然报Rank(A) p或者求解器返回 None甚至整个进程挂在 qp 调用上。根因通常是旧任务梯度之间出现了线性相关导致 G 矩阵退化Q 矩阵不满秩。线性相关在任务相似或者样本太少时特别容易出现比如你连续训两个都是自然图像分类的任务梯度方向高度相似。我的处理套路分三层。第一层是加一个微小的对角正则把 Q 换成 Q εIε 取 1e-6 到 1e-5相当于给问题加一点 L2 项让矩阵满秩绝大多数情况这就能把求解器哄好。第二层是求解失败时 fallback直接返回原始梯度 g_new这一次不做约束。有人担心 fallback 会破坏持续学习效果实测下来偶尔几次失败对整体精度的影响可以忽略因为约束是在绝大多数步上生效的。第三层是数据层面的如果失败特别频繁说明旧任务梯度太像了可以把每个任务的采样样本数量 b 调大一点增加梯度的多样性。还有个数值精度问题。如果你全程用 float32QP 求解器在参数维度很大、梯度数值范围跨度很大时容易报精度错误。我的做法是把送入求解器的数据全部转成 double算完再转回来。这个转换的显存开销可以接受因为 Q 和 p 的规模都很小只有任务数量级。4.2 内存和计算开销的控制GEM 的开销主要来自两块记忆样本的存储以及每个 batch 都要重算所有旧任务的梯度。存储这块好办每个任务存几百个样本十个任务也就几千张图普通显卡吃得下。真正压垮人的是梯度重算。前面算过这个开销随着任务数线性增长到了几十个任务时每个 batch 要跑几十次前向反向训练速度慢到没法用。我能给的实用优化有这么几条。其一控制旧任务采样数 b很多人舍不得觉得多点约束更稳其实 b10 和 b50 的效果差别很小但速度差 5 倍。其二用 A-GEM 的思路做近似把旧梯度先平均成一个向量约束变一条QP 退化成向量的方向判断速度快到几乎无感。其三把记忆样本预加载到显存别每次从硬盘读I/O 在 GEM 里也是隐性瓶颈尤其是记忆大了以后。其四如果只是为了验证方法不必跑完整的任务序列取前五个任务看趋势就够了。提示如果你的任务是图像类记忆样本建议存成 uint8用的时候再归一化成 float能省 4 倍存储空间对精度几乎没影响。4.3 任务顺序敏感与梯度噪声持续学习里有个老生常谈的现象叫任务顺序敏感同样的任务集合换个学习顺序最终精度能差十几个点。GEM 也躲不开这个问题而且它对顺序的敏感度和任务间的冲突程度直接相关。如果你把两个高度冲突的任务排在相邻位置投影约束会频繁触发新任务学得很吃力如果排在相隔较远的位置中间穿插了别的任务冲突可能被稀释掉。处理这个问题的实用办法是做顺序消融实验。别只跑一个任务顺序就下定论至少跑三到五个随机顺序看平均精度和方差。如果方差特别大说明你的方法对顺序不稳定这时候要回到超参上找原因往往是学习率太大或者记忆太小。我见过有人拿单次实验的最好结果去汇报被审稿人一个顺序消融就打回来了这个坑别踩。梯度噪声是另一个隐性问题。当 b 很小时旧任务梯度的估计噪声很大投影出来的方向可能一会儿约束强一会儿约束弱训练曲线会抖。缓解办法是增大 b或者用梯度累积的平滑技巧把几步的旧梯度平均一下再用。后者开销增加不多但对稳定性的提升明显是性价比很高的做法。4.4 常见问题速查表现象可能原因排查方向训练中途报求解器无解Q 矩阵不满秩加对角正则调大采样数旧任务精度依然掉得快学习率过大近似失效学习率降一半再试新任务学不动约束过强任务冲突大换任务顺序减少记忆约束训练速度极慢旧任务梯度重算开销减小 b改用 A-GEM 近似结果波动大梯度噪声或顺序敏感增大 b跑多个随机顺序显存溢出记忆样本太多存 uint8减少每任务样本数这张表我基本是拿自己的踩坑记录整理出来的放在手边随时查。我要额外强调一点GEM 的很多问题不是孤立的往往是几个因素一起作用。比如精度掉得快可能既是学习率的问题也是任务顺序的问题你需要一个个变量隔离出来测别一次改一堆参数那样你永远不知道是哪个起了作用。5. 关于 GEM 的取舍与后续演进讲到这里我得把话说得实在一点GEM 是一个思路漂亮、但工程上偏重的方法它不是万能药。它的优势在于约束的显式性和对正向迁移的天然支持效果在任务数不多十个以内、任务间存在一定相关性的场景下很稳。但它的短板同样明显每个 step 都要重算所有旧任务梯度任务数一大就没法用QP 求解的数值稳定性是个持续的心病对学习率敏感调参比普通训练费劲。所以实际项目里怎么选我的经验是分场景。如果任务只有三五个且你追求旧任务精度的下限保障GEM 值得一试它的约束能给你比较确定的防遗忘效果。如果任务几十上百个或者你根本没有精力调 QP那直接上 A-GEM 或者后来的经验回放类方法别死磕 GEM。A-GEM 把多个旧任务约束简化成一个平均梯度约束QP 退化成一维判断速度快几十倍精度损失通常在可接受范围内工程性价比高得多。这也是为什么现在很多持续学习的工程实现里跑的是 A-GEM 而不是原始 GEM。GEM 还启发了后面一系列工作。比如你在约束里引入样本重要性加权的思路就能演化成对关键参数保护更强的变体把静态记忆换成动态采样可以针对性地存那些容易被遗忘的样本把一阶近似换成二阶信息能缓解学习率敏感的问题但计算代价更高。你如果打算在 GEM 基础上做研究从这几个方向切入都比较有戏尤其是采样策略和近似阶数这两块工程上的可操作空间大。最后分享一个我压箱底的小体会。我刚开始复现 GEM 时总觉得是代码写错了因为精度上不去。折腾了半个月才发现问题出在我用了一个太大的学习率一阶近似根本没成立约束形同虚设。后来我把学习率降到原来的三分之一同时把每个任务的训练轮数翻倍整个曲线一下子就稳了旧任务精度保持得非常好。所以你要是也在为 GEM 效果不理想发愁不妨先别怀疑数学回头看看你的学习率和步长十有八九是这里的问题。
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进