1. 为什么持续学习绕不开 GEM 这道坎
做 Continual Learning 的人,早晚都会撞上 Gradient Episodic Memory(GEM)这个名字。它是 2017 年 NeurIPS 上的工作,作者是 David Lopez-Paz 和 Marc'Aurelio 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_k,d 是参数维度。这个问题有约束、有二次目标,正是标准的二次规划(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. 从零实现一个能跑的 GEM
3.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) 个样本。当 t=10、b=10 时,额外开销是 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, dim=0) # 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(), dtype=g_new.dtype, device=g_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 再给 cvxopt,float32 有时候会因为精度问题让求解器报错,这点很多人卡过。
还有个细节,如果旧任务梯度数量很多(比如超过 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:offset+numel].view_as(p)) offset += numel optimizer.step()有两处特别容易出 bug。一个是第 3 步里的model.zero_grad(),你在算旧任务梯度前必须把上一步的梯度清掉,否则不同任务的梯度会累加,投影出来的约束完全错误。另一个是第 5 步的偏移量管理,一定要保证参数顺序和拼梯度时的顺序完全一致,我见过有人因为模型里加了 BatchNorm 层,参数的遍历顺序和构建梯度时不一致,结果投影后的梯度错位,训练效果一塌糊涂。
3.5 参数配置与调参经验
给你一份我实测下来比较稳的起点配置,你可以在此基础上微调。
| 参数 | 建议值 | 说明 |
|---|---|---|
| 记忆大小 m | 200 到 500(每任务) | 任务越难越大,显存够就多存 |
| 旧任务采样数 b | 10 | 太小梯度噪声大,太大速度慢 |
| 学习率 | 0.01 到 0.05 | 比普通训练小,保证一阶近似成立 |
| batch 大小 | 10 到 30 | 单任务 batch,别设太大 |
| 优化器 | SGD 或 Adam | SGD 更贴合论文,Adam 也能跑 |
| 训练轮数 | 每任务 10 到 30 | GEM 收敛慢,要给足轮数 |
学习率这件事值得单独说。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 调大一点,增加梯度的多样性。
还有个数值精度问题。如果你全程用 float32,QP 求解器在参数维度很大、梯度数值范围跨度很大时,容易报精度错误。我的做法是把送入求解器的数据全部转成 double,算完再转回来。这个转换的显存开销可以接受,因为 Q 和 p 的规模都很小,只有任务数量级。
4.2 内存和计算开销的控制
GEM 的开销主要来自两块:记忆样本的存储,以及每个 batch 都要重算所有旧任务的梯度。存储这块好办,每个任务存几百个样本,十个任务也就几千张图,普通显卡吃得下。真正压垮人的是梯度重算。前面算过,这个开销随着任务数线性增长,到了几十个任务时,每个 batch 要跑几十次前向反向,训练速度慢到没法用。
我能给的实用优化有这么几条。其一,控制旧任务采样数 b,很多人舍不得,觉得多点约束更稳,其实 b=10 和 b=50 的效果差别很小,但速度差 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 效果不理想发愁,不妨先别怀疑数学,回头看看你的学习率和步长,十有八九是这里的问题。