GEM 持续学习实战:梯度投影解决灾难性遗忘 📅 发布时间:2026/9/16 18:06:46 👁 浏览次数: 一个在 MNIST 上做到 99.2% 的五分类器把剩下的五个类别按顺序继续喂给它等所有任务都学完再回头测第一个任务准确率只剩 31%。参数还在模型没坏但它几乎把最早学到的东西原封不动地覆盖掉了。这就是持续学习Continual Learning里最核心的那个问题灾难性遗忘。Gradient Episodic MemoryGEM是 2017 年 NeurIPS 上给出的一个经典解法思路出奇地干净——不改损失函数不加正则项只在更新方向跑偏的时候把它掰回一个安全的角度。它属于回放加约束这一派和 EWC 那种给权重拴弹簧的做法完全是两条路。这篇内容适合两类人看一是已经跑通基础分类或检测训练、正在被数据流式到来折磨的工程师二是正啃持续学习论文、准备复现经典 baseline 的同学。我会先把遗忘在梯度层面到底怎么发生讲透再把 GEM 的约束条件和那一步二次规划QP拆开然后给一份能直接跑的 PyTorch 骨架接着是我在项目里踩过的四个大坑最后是对比选型和检索关键词上的注意事项。1. 灾难性遗忘在梯度层面到底长什么样1.1 一个二维例子把覆盖这件事说透先抛开神经网络考虑只有两个参数的模型。任务 A 和数据分布 A 对应一组梯度方向任务 B 对应另一组。假设在某个参数点上任务 A 的梯度是 $g_A(1,0)$任务 B 的梯度是 $g_B(-1,0)$。它们的内积 $\langle g_A,g_B\rangle-1$是负的。现在按任务 B 的梯度走一步学习率 $\eta0.1$新参数 $\theta\theta-0.1g_B$。对任务 A 的损失做一阶展开$$L_A(\theta) \approx L_A(\theta) - \eta\langle g_B, g_A\rangle L_A(\theta) 0.1$$损失不但没降还涨了。走一百步任务 A 的损失就涨了 10 个单位。这就是覆盖的数学形态两个任务的梯度方向夹角超过 90 度时沿着其中一个走必然抬高另一个的损失。真实网络里情况会更糟。因为参数是共享的早期任务的解往往落在一个很窄的谷底任务 B 的梯度在共享层里和任务 A 的梯度几乎正交甚至反向而且网络越深、容量越大这种冲突越不容易被稀释反而更容易被整个抹平。这也是为什么在小网络上遗忘没那明显、换到 ResNet 上直接崩盘的原因之一。有一点很容易被忽略如果你刚好停在任务 A 的精确最优点上那 $g_A0$一阶分析会说这一步没伤害。但实际训练里我们从来不会精确停在最优点——随机梯度的噪声、有限步长、后期学习率衰减都让参数在最优附近抖动。一旦离开那个点$g_A$ 就不再是零约束才有意义。我见过有人拿最优点处约束是空的来质疑 GEM 的合理性其实是把分析点选错了。1.2 正则化、回放、参数隔离三条路的取舍持续学习的解法大致能装进三个抽屉理解它们的代价模型比记住名字重要得多。路线代表方法需要存旧数据额外内存随任务数增长主要短板正则化约束EWC、SI、MAS不需要线性增长每个任务一组 Fisher 对角任务数一多多个惩罚项互相打架回放/重演GEM、A-GEM、GSS、iCaRL需要可控固定总预算存样本有隐私和授权问题参数隔离PNN、PackNet、HAT不需要网络规模线性增长结构复杂部署时不好剪正则化那条路的思想是重要权重要少动用 Fisher 信息矩阵衡量重要性。听着优雅但实测里任务超过十个之后累计的惩罚项会把模型压得几乎学不动新任务你不得不在学得动和忘得少之间来回调系数。参数隔离那条路很干脆给每个任务分一块专属参数互不干扰。代价是模型体积线性膨胀端侧部署基本别想。GEM 选了回放这条路但它没有简单地把旧数据混进 batch 一起训——朴素回放只保证旧任务数据被看到不保证旧任务的损失不上升。GEM 在回放的基础上加了一层硬约束这一步更新之后所有旧任务在内存样本上的损失一阶近似不增加。硬约束和软惩罚的区别是 GEM 最值得学的设计。2. 从旧任务损失不能涨到梯度夹角必须小于 90 度2.1 一阶泰勒展开给出的充分条件设当前参数为 $\theta$当前任务 $t$ 的梯度 $g\nabla L_t(\theta)$历史任务 $k$ 的梯度 $g_k\nabla L_k(\theta)$。沿 $g$ 走一步得到 $\theta-\eta g$对历史任务的损失做一阶展开$$L_k(\theta-\eta g) \approx L_k(\theta) - \eta\langle g, g_k\rangle$$只要学习率 $\eta0$想让 $L_k$ 不上升只需要 $\langle g,g_k\rangle\ge 0$。对所有历史任务 $k1,\dots,t-1$ 同时成立就得到 GEM 的核心约束$$\langle g, g_k\rangle \ge 0,\quad \forall kt$$几何上每个历史任务定义一个安全半空间所有半空间的交集是一个凸锥。落在锥里的方向随便走落在外面的方向必须先投影回来。这就是为什么 GEM 的插图总是画一个扇形区域。这里有个工程上很实用的推论约束只依赖方向不依赖步长。所以你不能靠调小学习率来绕过约束——学习率调小只是让每步的伤害小一点方向冲突依旧存在任务多了照样积累成大遗忘。真正解决问题的只有改变方向。2.2 约束不满足时投影到可行锥如果当前梯度 $g$ 不在锥里我们想找一个最接近它的可行方向$$\tilde{g} \arg\min_{z}\ \tfrac{1}{2}|z-g|^2 \quad \text{s.t.}\quad Mz \ge 0$$其中 $M$ 是把所有历史任务的梯度按行堆起来的矩阵形状是 $(t-1)\times d$$d$ 是参数量。直接对 $z$ 求解是不现实的——$d$ 动辄几百万。GEM 的关键一步是转到对偶空间。构造拉格朗日函数令 $zgM^\top\alpha$$\alpha\ge 0$问题变成$$\alpha^\star \arg\min_{\alpha\ge 0}\ \tfrac{1}{2}\alpha^\top MM^\top\alpha \alpha^\top Mg$$$$\tilde{g} g M^\top\alpha^\star$$这一步是整个方法最漂亮的地方。变量数从参数量 $d$ 降到了历史任务数 $t-1$。二十个任务就是二十个变量的 QP随便什么求解器都能秒解。你在参数空间里做投影那是几百万维的凸优化搬到任务空间它变成几十维的小问题。$\alpha$ 的取值还有明确的解释$\alpha_i0$ 说明第 $i$ 个历史任务的约束是紧的正在起作用$\alpha_i0$ 说明那个任务本来就安全不需要管它。调试的时候把 $\alpha$ 打出来看一眼能立刻知道是哪几个任务在跟当前任务抢方向。2.3 这个保证有多硬必须说清楚 GEM 的保证是一阶、局部的。它假设损失曲面在当前点附近近似线性梯度在内存样本上的估计能代表整个旧任务分布优化器直接沿着 $\tilde{g}$ 更新参数。这三条任何一条不成立约束的实际效果都会打折扣。原论文里报告 BWT后向迁移始终非负这个结论在 MNIST 置换任务和 CIFAR 多任务上都成立但那是配合固定的小学习率和 SGD 得到的。换成 Adam、加大 batch、或者用很激进的余弦退火我都见过 BWT 变负的情况。别把论文里的趋势当成定理。3. 手写一版 GEM缓冲区、逐任务梯度与 QP 落地3.1 内存缓冲区的组织方式GEM 的内存是按任务分桶的不是一个大池子。每个任务一个桶桶内有固定容量 $m$总预算 $M m\times T$。这个设计不是为了好看而是因为计算 $g_k$ 时需要明确知道这批样本属于哪个历史任务。采样策略我从简到繁试过三种纯随机替换桶满了之后随机决定替不替换、替换谁。实现最简单分布近似均匀我大部分情况下用它。蓄水池采样严格保证桶内是流式数据的一个均匀子集。数据分布随时间漂移比如新品类的特征分布和老品类差异很大时更稳。类别均衡 随机每个类别分配固定配额桶内再随机。类别不平衡的数据集上这个比纯随机的效果好一大截代价是要多存一个类别计数。容量怎么定论文里 CIFAR 用了每任务 200 个样本MNIST 置换用的是 256。我的经验法则是总内存控制在训练集总量的 1% 到 5%然后按任务数均分。如果任务数超过 50每任务分不到 50 个样本时GEM 的约束估计会明显失真这时候要么改用 A-GEM 那种更粗但更稳的约束要么干脆上生成式回放。还有一个细节很多人漏掉存进内存的样本要不要做数据增强。我的做法是存原图采样出来之后再在线增强。这样一能省内存二能让每次算 $g_k$ 时看到的样本略有不同相当于给约束加了一点点随机性反而比固定样本更不容易过拟合到那几个特定样本上。3.2 单次迭代的完整数据流搞清楚顺序很重要写错了梯度会互相污染从当前任务取一个 batch前向 反向得到 $g$。注意用torch.autograd.grad而不是backward前者不会累积到.grad上。遍历所有见过的历史任务每个任务从对应桶里采一个 batch前向 反向得到 $g_k$。这一步是 GEM 的主要开销任务数 $T$ 就意味着每步要做 $T$ 次反向传播。把 $g_k$ 堆成矩阵 $M$解 QP 得到 $\alpha^\star$算出 $\tilde{g}$。把 $\tilde{g}$ 手动写回每个参数的.grad再调optimizer.step()。把当前 batch 的样本按策略塞进当前任务的桶里。第 4 步是最容易写错的地方。你必须先zero_grad()再把投影后的梯度切片赋值给每个参数而不是让优化器自己算。切片的时候要按参数在列表里的顺序累加偏移量顺序错一位梯度就整个错位了。3.3 PyTorch 骨架代码下面这段是我实际在用的结构删掉了项目里的日志和分布式部分核心逻辑保留import numpy as np import torch from quadprog import solve_qp def flat_grad(loss, params, retain_graphFalse): 把 loss 对各参数的梯度拍平成一维向量。 grads torch.autograd.grad( loss, params, retain_graphretain_graph, allow_unusedTrue ) return torch.cat([ (g if g is not None else torch.zeros_like(p)).reshape(-1) for g, p in zip(grads, params) ]) class EpisodicMemory: 每个任务一个独立桶桶满后随机替换。 def __init__(self, n_tasks, per_task): self.buckets [[] for _ in range(n_tasks)] self.per_task per_task def add(self, task_id, x, y): bucket self.buckets[task_id] for i in range(x.size(0)): item (x[i].detach().cpu(), y[i].detach().cpu()) if len(bucket) self.per_task: bucket.append(item) else: # 蓄水池采样的简化版 j np.random.randint(len(bucket) 1) if j len(bucket): bucket[j] item def sample(self, task_id, batch_size, device): bucket self.buckets[task_id] idx np.random.randint(0, len(bucket), sizebatch_size) xs torch.stack([bucket[i][0] for i in idx]).to(device) ys torch.stack([bucket[i][1] for i in idx]).to(device) return xs, ys def gem_update(model, optimizer, params, x, y, memory, task_id, criterion, device, mem_batch64, eps1e-4): # 1) 当前任务的原始梯度 optimizer.zero_grad() loss_cur criterion(model(x), y) g flat_grad(loss_cur, params) # 2) 历史任务在各自内存桶上的梯度 g_old [] for k in range(task_id): if len(memory.buckets[k]) 0: continue bx, by memory.sample(k, mem_batch, device) optimizer.zero_grad() loss_k criterion(model(bx), by) g_old.append(flat_grad(loss_k, params)) # 3) 无历史任务或直接解 QP 求投影 if not g_old: g_tilde g else: M torch.stack(g_old) # (t-1, d) MMt (M M.T).double().cpu().numpy() Mg (M g).double().cpu().numpy() MMt MMt eps * np.eye(MMt.shape[0]) # 抖动保证正定 # solve_qp 求解 min 0.5 aPa - qa, s.t. a 0 alpha solve_qp(MMt, -Mg, np.eye(len(Mg)), np.zeros(len(Mg)), 0)[0] alpha torch.as_tensor(alpha, dtypeg.dtype, deviceg.device) g_tilde g alpha M # 4) 手动把投影后的梯度灌回参数 optimizer.zero_grad() offset 0 for p in params: n p.numel() p.grad g_tilde[offset:offset n].view_as(p).clone() offset n optimizer.step() memory.add(task_id, x, y)有一个地方值得单独提醒(M M.T)在 float32 下很容易出现数值问题尤其是参数量大、梯度尺度差异大的时候。我在项目里统一用 double 算这一步再转回 float32稳定性提升很明显代价只有几毫秒。3.4 QP 求解器怎么挑求解器依赖优点缺点quadprog纯 Python 包轻快、接口简单、结果稳定只支持稠密矩阵任务数过百会慢cvxopt老牌库论文原版用它支持各种约束形式安装偶尔出问题API 啰嗦qpth基于 PyTorch可微能端到端反传依赖多小问题上没必要自写投影梯度无零依赖精度差需要调迭代次数我的默认选择是quadprog。二十到五十个任务的规模下单次求解通常在几十微秒相对于一次反向传播完全可以忽略。真正的时间瓶颈是第 2 步里那 $t-1$ 次反向传播不是 QP。如果实在不想引入依赖可以用迭代投影反复对违反的约束做单约束投影跑个十几轮。精度会差一点但在约束大多满足的情况下也就是 $\alpha$ 大部分为零实际效果差别不大。这个技巧在部署环境受限时很好用。4. 实测里最容易翻车的四个地方4.1 优化器把投影悄悄吃掉了这个坑我踩得最久。GEM 的保证是关于原始梯度的但如果你用 Adam实际更新方向是梯度的一阶矩和二阶矩的归一化结果和 $\tilde{g}$ 已经不是一回事了。动量会记住之前几步的方向把这一步的投影效果稀释掉。我试过在置换 MNIST 上把 SGD 换成 AdamBWT 直接从 0.02 掉到 −0.09。排查链路是这样的先怀疑内存容量不够把每任务样本从 100 加到 400没改善再怀疑梯度尺度加了归一化好了一点但不明显最后把优化器换回带动量的 SGDBWT 立刻回正。所以我的建议是上 GEM 就用 SGD 动量0.9学习率用 0.01 到 0.1配余弦退火。如果非要用 Adam那就把投影作用在优化器输出的实际更新量上而不是原始梯度上——代价是要改优化器内部逻辑工程量不小。4.2 BatchNorm 的滑动统计量在偷偷漂移这个问题比梯度冲突更隐蔽。BatchNorm 在训练时会更新 running_mean 和 running_var。当你从任务 A 切到任务 B前向传播经过 BN 层时这些统计量会被任务 B 的数据一步步改写。等到回头测任务 A即使卷积核的权重一点没变BN 的统计量已经完全不匹配了准确率照样掉。我第一次遇到的时候非常困惑明明梯度约束都满足了为什么旧任务还是掉后来把 BN 的 running 统计量冻结住掉幅立刻从 40% 缩到 8%。三种处理方式按推荐度排序换成 GroupNorm 或 LayerNorm彻底没有跨任务统计量问题或者把 BN 的 momentum 设成 0、训练时冻结 running 统计量只更新可学习的 affine 参数再或者给每个任务存一份 BN 统计量测试时按任务切换。最后一种最灵活但最麻烦任务数多的时候管理成本很高。4.3 梯度尺度不一致让 QP 偏向某个任务不同任务的梯度范数可能差好几个数量级。比如一个任务的分类头刚初始化、loss 很大梯度范数可能是另一个任务的几十倍。QP 在解的时候约束矩阵 $MM^\top$ 的对角元素就是各梯度的平方范数尺度大的任务会主导整个求解$\alpha$ 几乎全分给它。我处理的办法是在堆矩阵之前把每个 $g_k$ 归一化到单位范数$g$ 保持原尺度。这样约束就变成了方向约束而不是幅度约束符合 GEM 的原始直觉。归一化之后如果发现约束经常被违反说明方向冲突是真实存在的不是尺度造成的假象。还有一个更稳的做法是给每个任务的约束加一个松弛系数 $\gamma_k$把约束改成 $\langle g,g_k\rangle \ge -\gamma_k$。$\gamma_k$ 可以按任务的重要程度或者数据量来定。这就是软约束版本的 GEM实践中比硬约束宽容得多代价是引入了需要调的参数。4.4 任务数一多QP 和反向传播一起拖慢理论上 GEM 每步要 $t-1$ 次反向传播。十个任务就是十一次二十个任务就是二十一次。训练时间基本按任务数线性增长这一点在论文里说得比较轻描淡写实际跑起来很痛。缓解手段有几个我都试过降低内存采样的 batch 大小。算 $g_k$ 本来就是个估计用 32 个样本和用 128 个样本对结果影响很小但时间差四倍。缓存梯度。如果历史任务的梯度在几步之内变化不大可以每 N 步才重算一次。我试过 N5速度提升约 3 倍BWT 只掉了 0.01。只对最近访问过的任务施加约束。这个改动有点激进等于放弃了长期保护但在任务之间相似度高的场景下效果可以接受。换成 A-GEM。这是最直接的办法后面细说。5. 选型对照GEM、A-GEM、EWC 与朴素回放该怎么挑5.1 四个方法摆在一张表上维度朴素微调EWCGEMA-GEM是否存旧数据否否是是每步反向传播次数11t任务数2额外内存无O(t·d)O(M·样本大小)同 GEM遗忘抑制强度无中强中到强对学习率敏感度低中高中任务数扩展性—差中好实现复杂度低低中高低EWC 的额外内存和任务数线性相关每个任务存一份 Fisher 对角矩阵参数量大的模型上这个开销不小。GEM 的内存是样本只要总预算固定任务数增加不会让内存膨胀只是每任务分的样本变少。5.2 A-GEM两行代码把 QP 干掉A-GEM 的核心简化是不去逐个约束所有历史任务而是把所有历史任务的梯度取平均得到一个参考梯度 $g_{ref}$然后只检查一个约束 $\langle g,g_{ref}\rangle\ge 0$。如果违反就把 $g$ 投影到 $g_{ref}$ 的正交方向上$$\tilde{g} g - \frac{\langle g, g_{ref}\rangle}{|g_{ref}|^2}g_{ref}$$就这三行公式把每步的反向传播次数从 $t$ 降到 2一次当前任务一次内存采样把 QP 求解整个省掉。我在 CIFAR 十任务上对比过A-GEM 的最终平均准确率比 GEM 低大约 1 到 2 个点但训练速度快 5 倍以上。任务数超过 20 之后A-GEM 的性价比明显更高。什么时候坚持用完整 GEM我的判断标准是任务数少于 10、任务之间冲突明显、而且你有充足的训练时间预算。这种情况下 GEM 的逐任务约束确实能多榨出一点精度。5.3 什么场景下我会直接放弃 GEM有几种情况我会建议绕开这类方法数据不能留存的时候。医疗影像、金融风控、用户隐私数据你把样本存到内存里可能就是合规问题。这种情况下只能走正则化路线或者用生成式回放训一个生成模型用生成的样本代替真实样本。任务边界模糊的时候。GEM 假设任务边界清晰、训练时知道当前是任务几。如果数据是连续流、任务之间的切换是渐变的内存分桶的逻辑就不成立了。这时候要用任务无关的聚类或者在线分割方法。需要类增量而不是任务增量的时候。GEM 默认测试时知道任务 ID分类头可以按任务选。如果测试时不知道任务 ID、要在所有类别里做单头分类那分类头的 logits 会跨任务漂移需要额外加 logit 校准或原型分类器。5.4 指标别只看平均准确率两个指标一定要同时看ACC所有任务训练完后在所有任务测试集上的平均准确率。反映最终水平。BWT$\frac{1}{T-1}\sum_{i1}^{T-1}(R_{T,i}-R_{i,i})$即学完所有任务后任务 $i$ 的准确率减去刚学完任务 $i$ 时的准确率。负值代表遗忘正值代表后学到的任务对前面有正向迁移。我见过不少复现报告只给 ACC结果 ACC 看起来还行BWT 其实是 −0.3说明模型是靠后面任务的表现把平均值撑起来的前面的任务已经废了。两个指标必须一起报。计算 BWT 需要在每个任务训练结束时都跑一遍所有已见任务的测试集这是 $O(T^2)$ 次评估。任务数不多的时候无所谓任务多了要记得把这个开销算进实验预算里。6. 检索与复现的最后清单6.1 GEM 这个词在搜索里的重名问题搜资料的时候有个很实际的干扰GEM 这个缩写在工业界指代的是半导体设备通信标准 SECS/GEM和持续学习毫无关系。我用搜索引擎找代码实现的时候前几页经常混进设备通信协议的内容。有效的做法是在关键词里加上限定词比如gradient episodic memory continual learning、GEM NeurIPS 2017、gem pytorch continual learning或者直接搜论文标题。中文资料里梯度情景记忆这个译名偶尔也会出现但用的人不多主要还是靠英文关键词。另一个小坑是有些开源实现把 GEM 和 A-GEM 放在同一个仓库里类名都叫GEM或者Gem配置里靠一个布尔开关切换。你要是没注意可能会以为自己在跑完整版其实跑的是平均梯度版实验结论就完全对不上了。跑之前一定去代码里确认一下有没有solve_qp或者cvxopt的调用。6.2 复现时按这个顺序检查下面这张表是我给自己团队新人准备的检查顺序基本能覆盖 90% 的复现失败现象优先检查项旧任务准确率掉得比朴素微调还多投影后的梯度有没有正确写回.grad参数切片偏移对不对BWT 一直在 0 附近不动约束是不是从没被违反过打印一下 $\alpha$ 看是不是全零训练到第五六个任务突然崩内存桶是不是写空了采样数小于 batch size 会报错或采重复换到自己的数据集效果差任务间相似度太低内存样本不具代表性考虑增大每任务容量速度慢到不可接受内存采样 batch 大小、是否缓存梯度、优化器是不是 Adam结果方差特别大学习率太高、种子太少跑三次以上取平均参数起点我一般这么设内存总量取训练集的 2%每任务样本数 总量 / 任务数内存采样 batch 64学习率 SGD 0.01 到 0.1 按任务难度调动量 0.9权重衰减 0 或 1e-4QP 的抖动项 1e-4。这些不是最优值但是一个不容易翻车的起点。我个人到现在还保留的一个习惯是跑 GEM 之前先把朴素微调和联合训练所有任务数据一起训两条基线都跑出来。朴素微调是下界联合训练是上界你的 GEM 结果落在什么位置一眼就知道是方法问题还是实现问题。很多时候我发现代码有 bug就是因为最终的 ACC 居然比朴素微调还低——这在那套设置下是不可能的。