AI持续学习:对抗灾难性遗忘的工程实践
引言
模型上线不是终点,而是学习的起点。推荐系统每天有新用户行为,风控模型每月面对新的欺诈手法,语音助手要不断学新方言。理想情况是模型像人一样持续吸收新知识,但现实很骨感:用新数据直接微调,旧任务上的表现会断崖式下跌——这就是灾难性遗忘(Catastrophic Forgetting)。持续学习(Continual Learning)研究的就是如何让模型"学而时习之",在学新任务时不丢掉旧能力。本文从遗忘的机理讲起,梳理三大技术路线,并给出可在生产中落地的工程方案。
灾难性遗忘是怎么发生的
神经网络的参数是共享的。任务A学完后,参数落在一个对A友好的区域;用任务B的数据继续训练,梯度会毫不犹豫地拉动参数走向对B友好的区域,如果两个区域不重叠,A的性能就毁了。问题的根源在于:梯度下降只关心当前损失,完全不记得参数对旧任务有多重要。
还有一个更隐蔽的因素:表征漂移。即便输出层做了保护,backbone的权重变化会让旧数据的特征表示失效,下游的一切统计都跟着作废。所以持续学习必须同时解决"参数怎么走"和"特征怎么稳"两个问题。
需要区分几个相近概念:多任务学习是一次性学所有任务,数据都在手上;迁移学习是学完A就不管A了,只追求B的效果;持续学习是任务按顺序到来、旧数据不可得或只能少量保留,且要求旧任务性能不掉。第三种设定最苛刻,也最贴近生产。
三大技术路线
正则化方法:给损失函数加惩罚项,让"对旧任务重要的参数"不轻易动。代表作EWC(Elastic Weight Consolidation)用Fisher信息矩阵估计每个参数对旧任务的重要性,重要性越高,偏移惩罚越大。MAS用输出对参数的敏感度替代Fisher,思路类似。LwF(Learning without Forgetting)则不加参数惩罚,而是用旧模型在新数据上的输出做知识蒸馏,约束新模型的行为。这类方法不占额外存储,但任务多了之后约束会互相打架。
回放方法:最直接——留一小部分旧数据(或生成伪样本),训练新任务时混进去一起学。iCaRL用"最接近类均值"的样本构成核心集;GEM用旧任务梯度约束新任务的梯度方向,保证旧任务损失不增;DER(Dark Experience Replay)连旧模型的logits一起存,蒸馏加回放双管齐下,效果常年霸榜。回放方法简单粗暴但有效,代价是存储和隐私——某些行业根本不允许保留原始数据。
结构方法:给每个任务分配专属参数。PackNet通过剪枝释放冗余容量,每个任务占用一部分神经元;Progressive Network为新任务新增一列网络,彻底不干扰旧任务。隔离效果最好,但参数量随任务数膨胀,推理部署也麻烦。
工程实战:EWC与回放的组合方案
实际生产中,单一方法往往不够,通常组合使用。下面是一个EWC的核心实现,配上经验回放就是工业界常用的baseline:
import torch import torch.nn as nn class EWC: """记录旧任务的Fisher信息和最优参数,训练新任务时施加惩罚""" def __init__(self, model, dataloader, device, sample_size=200): self.device = device self.params = {n: p.clone().detach() for n, p in model.named_parameters()} self.fisher = self._compute_fisher(model, dataloader, sample_size) def _compute_fisher(self, model, dataloader, sample_size): fisher = {n: torch.zeros_like(p) for n, p in model.named_parameters()} model.eval() count = 0 for x, y in dataloader: if count >= sample_size: break model.zero_grad() out = model(x.to(self.device)) loss = nn.functional.cross_entropy(out, y.to(self.device)) loss.backward() for n, p in model.named_parameters(): fisher[n] += p.grad.detach() ** 2 count += x.size(0) return {n: f / count for n, f in fisher.items()} def penalty(self, model): loss = 0.0 for n, p in model.named_parameters(): loss += (self.fisher[n] * (p - self.params[n]) ** 2).sum() return loss # 训练新任务时: # total_loss = new_task_loss + lambda_ewc * ewc.penalty(model) # lambda_ewc 通常在 1e2 ~ 1e4 之间调 # 值越大越保旧任务,但新任务越难学进去回放部分只需维护一个固定大小的buffer,新任务训练时按1:3到1:1的比例混入旧样本。buffer更新策略推荐水库采样(Reservoir Sampling),保证每个历史样本被选中的概率均等,避免buffer被近期数据占满。
上线前还有几个工程细节:任务切换点要做全量回归评测,旧任务性能下降超过阈值就报警回滚;Fisher矩阵和buffer要跟模型一起做版本管理;如果数据合规不允许存原始样本,可以降级为只存特征或logits。
