AI持续学习:对抗灾难性遗忘的工程实践 引言模型上线不是终点而是学习的起点。推荐系统每天有新用户行为风控模型每月面对新的欺诈手法语音助手要不断学新方言。理想情况是模型像人一样持续吸收新知识但现实很骨感用新数据直接微调旧任务上的表现会断崖式下跌——这就是灾难性遗忘Catastrophic Forgetting。持续学习Continual Learning研究的就是如何让模型学而时习之在学新任务时不丢掉旧能力。本文从遗忘的机理讲起梳理三大技术路线并给出可在生产中落地的工程方案。灾难性遗忘是怎么发生的神经网络的参数是共享的。任务A学完后参数落在一个对A友好的区域用任务B的数据继续训练梯度会毫不犹豫地拉动参数走向对B友好的区域如果两个区域不重叠A的性能就毁了。问题的根源在于梯度下降只关心当前损失完全不记得参数对旧任务有多重要。还有一个更隐蔽的因素表征漂移。即便输出层做了保护backbone的权重变化会让旧数据的特征表示失效下游的一切统计都跟着作废。所以持续学习必须同时解决参数怎么走和特征怎么稳两个问题。需要区分几个相近概念多任务学习是一次性学所有任务数据都在手上迁移学习是学完A就不管A了只追求B的效果持续学习是任务按顺序到来、旧数据不可得或只能少量保留且要求旧任务性能不掉。第三种设定最苛刻也最贴近生产。三大技术路线正则化方法给损失函数加惩罚项让对旧任务重要的参数不轻易动。代表作EWCElastic Weight Consolidation用Fisher信息矩阵估计每个参数对旧任务的重要性重要性越高偏移惩罚越大。MAS用输出对参数的敏感度替代Fisher思路类似。LwFLearning without Forgetting则不加参数惩罚而是用旧模型在新数据上的输出做知识蒸馏约束新模型的行为。这类方法不占额外存储但任务多了之后约束会互相打架。回放方法最直接——留一小部分旧数据或生成伪样本训练新任务时混进去一起学。iCaRL用最接近类均值的样本构成核心集GEM用旧任务梯度约束新任务的梯度方向保证旧任务损失不增DERDark Experience Replay连旧模型的logits一起存蒸馏加回放双管齐下效果常年霸榜。回放方法简单粗暴但有效代价是存储和隐私——某些行业根本不允许保留原始数据。结构方法给每个任务分配专属参数。PackNet通过剪枝释放冗余容量每个任务占用一部分神经元Progressive Network为新任务新增一列网络彻底不干扰旧任务。隔离效果最好但参数量随任务数膨胀推理部署也麻烦。工程实战EWC与回放的组合方案实际生产中单一方法往往不够通常组合使用。下面是一个EWC的核心实现配上经验回放就是工业界常用的baselineimport torch import torch.nn as nn class EWC: 记录旧任务的Fisher信息和最优参数训练新任务时施加惩罚 def __init__(self, model, dataloader, device, sample_size200): 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。大模型时代参数高效微