持续学习Continual Learning不是一个全新的概念但在真实业务里它的必要性越来越明显。推荐系统每天都会收到新行为数据对话机器人要持续添加新的技能包质检模型要适配新产线、新材质甚至同一家公司的不同客户都会产生不同的数据分布。如果每次都把所有历史数据捞出来全量重训训练链路会越来越慢存储和算力成本也会持续上升如果只拿新数据做普通增量训练模型又会把已经学会的旧知识忘掉。如何在不断到来的数据上学习同时尽量保留旧本领正是持续学习要解决的问题。标题中的“in Transition”可以理解成两层含义持续学习研究正从固定基准实验过渡到真实数据流和工程系统模型训练也从“训练一次部署很久”过渡到“持续更新、持续验证、持续回滚”。很多综述会从理论、方法和应用三个维度组织这个领域这篇文章也按这条线索展开。1. 先理解持续学习要解决的核心难题灾难性遗忘1.1 灾难性遗忘的根源是共享参数被覆盖先看一个最直观的场景。假设模型先学任务 A在任务 A 上达到不错的准确率然后只拿任务 B 的数据继续训练。训练任务 B 时反向传播只计算任务 B 的损失对参数的梯度优化器也只会沿着降低任务 B 损失的方向更新参数。这个更新方向并不会主动维护任务 A 的决策边界所以共享参数会被逐步推向更适合任务 B、却可能破坏任务 A 特征表示的方向。这就是“灾难性遗忘”的本质神经网络把多个任务的知识存储在同一个参数空间里新任务的梯度更新覆盖了旧任务依赖的参数位置。旧任务的训练数据不再参与损失计算模型没有动力保持旧任务上的输出行为因此旧任务准确率会快速下降。如果把模型训练看成一个连续优化问题那么持续学习就是在没有完整数据访问权限的前提下寻找一个能同时满足多个任务损失的参数解。这个解不一定存在一个固定点因为新任务的数据分布、标签空间和损失曲面都会变化。所以持续学习不是简单地把“批量训练”改成“多跑几次训练”而是要重新设计损失函数、数据管理方式和模型结构。这里有一个容易混淆的点模型参数共享是灾难性遗忘的根源但完全隔离参数并不是唯一解法。很多持续学习算法不会为每个任务开一套独立参数而是想办法让新旧任务共享大部分特征同时用约束、回放或动态扩展来保护关键参数。具体选哪种取决于任务形态和数据约束。1.2 三种任务形态任务增量、领域增量、类别增量持续学习论文里经常出现三类任务设定它们的难度和应用场景差别很大。如果不先分清任务形态就很容易把实验结果理解错也容易在工程选型时做出错误判断。任务形态输入分布输出空间推理时是否知道任务 ID典型例子核心难点任务增量学习Task-Incremental每个任务有自己的数据分布每个任务有独立标签空间知道先学习数字 0-1再学习数字 2-3多个任务头共存共享特征不能漂移太严重领域增量学习Domain-Incremental输入分布不断变化输出类别保持不变不知道同一种商品在不同光照、角度、背景下的分类数据漂移导致特征分布偏移但标签含义不变类别增量学习Class-Incremental新类别不断出现输出空间不断扩展不知道相机不断识别新物种推理时要区分所有已见类别最容易遗忘三种形态中类别增量通常被认为最难因为模型不仅要知道“这个样本属于哪个新类别”还要在所有已经见过的类别之间做判别。任务增量由于推理时能拿到任务 ID可以给每个任务单独分配一个分类头难度会低一些。工程中很多场景并不是严格的某一种形态。比如工业质检先在某条产线训练后来换到另一条产线输入分布变了但缺陷类别还是那几类这更接近领域增量如果缺陷类型本身也在增加那就往类别增量靠。定义清楚任务形态才能决定评估指标和模型结构。1.3 Transition从一次性训练过渡到流式学习传统机器学习项目通常遵循“收集数据 - 清洗 - 训练 - 评估 - 发布 - 进入维护期”的流程。模型发布后如果数据分布变化常见的做法是把新旧数据合并重新训练一次然后发布新版本。全量重训在数据量可控时是可靠的甚至是最稳妥的方案因为它不依赖任何增量学习技巧。但真实业务里有三个问题会让全量重训越来越吃力数据累积速度超过训练资源扩张速度。部分旧数据因为隐私、版权、存储策略无法长期保留。业务要求模型尽快适应新增事件比如新促销规则、新品类、新故障模式。持续学习想要替代的并不是一次完整的全量重训而是那些“不能全量重训”或“全量重训成本过高”的场景。它让模型从“训练一次、部署很久”过渡到“持续接收数据、持续更新、持续验证”。这种训练方式的转变会带来新的工程问题数据版本怎么管理、模型如何回滚、旧能力如何监控、更新频率如何设计。这些内容会在后面的工程章节展开。因此“Continual Learning in Transition”可以理解成一次训练理念的过渡从“静态模型”走向“流式模型”从“以数据集为中心”走向“以任务序列为中心”。2. 持续学习的主流方法不是单点方案而是一组取舍持续学习方法没有绝对最优解只有针对不同约束的选择。按是否存储旧数据、是否扩展模型参数、如何避免遗忘可以分成三大类。2.1 基于正则化的方法在损失函数上做约束正则化方法的基本思路是学习新任务时不允许参数在旧任务重要的方向上有太大移动。典型代表是 EWCElastic Weight Consolidation。它的损失函数可以写成L(θ) L_new_task(θ) λ / 2 * Σ_i F_i * (θ_i - θ_old_i)^2其中θ_old是旧任务训练结束后保存的参数F_i是 Fisher 信息矩阵对角近似表示参数θ_i对旧任务的重要性。Fisher 越大说明参数越重要新任务更新时就越不要移动它。EWC 的优点是不需要保存旧数据只需要保存旧任务上计算出的 Fisher 对角向量和模型参数。缺点是当任务数量很多时累积约束会越来越复杂且 Fisher 只是局部重要性估计不同任务之间可能出现约束冲突。同类方法还有 LwFLearning without Forgetting。LwF 不保存旧样本而是把旧模型在新样本上的输出作为软标签让新模型尽量保持旧任务的输出分布。这种思路在数据敏感场景里很有用但依赖旧模型输出的质量如果旧模型本身在新样本上不稳定效果会打折扣。2.2 基于回放的方法用旧样本或生成样本对抗遗忘回放方法的核心是既然遗忘是因为旧数据没有参与训练那就尽量让旧知识以某种方式重新参与训练。最简单的做法是维护一个内存缓冲区训练新任务时混合采样一部分旧样本。这类方法叫 Experience Replay。回放为什么有效因为它最接近联合训练。如果每个 batch 里既包含新任务样本又包含旧任务样本模型就能在优化新损失的同时持续看到旧任务的数据分布。相比纯正则化方法回放通常更容易抑制遗忘尤其是类别增量场景。回放方法的代价也很直接内存开销缓冲区越大保留的旧知识越完整但存储成本越高。训练开销每个 step 要额外计算旧样本的损失。隐私风险如果业务规定旧数据不能长期保存直接保存样本会违规。缓冲区管理固定容量下哪些样本该留、哪些该淘汰本身就是优化问题。为了解决隐私问题很多研究尝试用生成模型生成伪样本替代真实样本。这类方法的效果依赖生成模型质量工程复杂度更高。2.3 基于动态架构的方法为每个任务分配独立参数动态架构方法不强行让所有任务共享同一套参数而是在新任务到来时扩展网络结构或者用掩码隔离不同任务使用的参数。比如 Progressive Networks 会为新任务增加新的特征通路PackNet 通过剪枝和掩码把网络分成多个任务区域。这类方法可以有效避免参数覆盖因为新任务不再直接改写旧任务的关键参数。代价是模型体积会随任务数量增长推理时可能还需要根据任务 ID 选择不同的网络路径或掩码部署复杂度会上升。在边缘设备和推理延迟敏感的系统中动态架构要谨慎使用。每增加一个任务模型体积都会变大如果任务数量很多最终可能比全量重训一个统一模型还要昂贵。2.4 方法对比与选型参考方法类别是否需要旧数据是否扩展参数典型算法适合场景主要成本正则化不需要不扩展EWC、LwF数据隐私要求高、任务数量有限约束累积任务多时可能效果下降回放需要缓冲区不扩展Experience Replay、iCaRL类别增量、数据可以留存存储和训练成本增加动态架构通常不需要扩展Progressive Networks、PackNet任务边界清晰、允许模型变大模型体积和推理复杂度上升实际项目中回放方法更容易获得稳定效果因为它最接近联合训练的近似。正则化方法使用简单但要注意参数约束强度。动态架构适合任务边界清楚、系统能容忍模型膨胀的场景而不适合任务数量无限增长的服务端模型。3. 从理论到落地先建立一套可复现的评估环境持续学习领域的一个常见问题是不同论文之间的实验设置差异很大直接比较结果很容易失真。为了真正理解一个方法是否有效应该自己把任务序列、评估指标和随机种子固定下来再对比不同方法。3.1 使用 Python 和 PyTorch 搭建最小实验环境下面的示例使用 PyTorch 实现一个最小的任务增量学习流程。学习环境只需要 CPU 就能跑通不需要一开始就准备 GPU。建议 Python 版本 3.9 以上PyTorch 版本 2.x 以上。如果当前环境版本不同落地前要确认 API 兼容性。python -m pip install torch torchvision numpy如果只使用 CPU 环境可以按 PyTorch 官方提供的 CPU 安装命令安装。版本号会随发布时间变化不建议直接复制别人的固定版本号应该以当前安装源的可用版本为准。需要准备的数据集是 MNIST。MNIST 虽然简单但非常适合观察灾难性遗忘因为普通网络在单个任务上很容易收敛搬到任务序列后能明显看到旧任务准确率回落。3.2 用经典基准数据集构造任务序列这里把 MNIST 按类别拆成三个任务任务 1 学习数字 0 和 1任务 2 学习数字 2 和 3任务 3 学习数字 4 和 5。每个任务内部是二分类。from torch.utils.data import DataLoader, Subset from torchvision import datasets, transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) def get_mnist_task(classes, trainTrue): dataset datasets.MNIST( root./data, traintrain, downloadTrue, transformtransform ) labels set(classes) indices [i for i, (_, y) in enumerate(dataset) if y in labels] return Subset(dataset, indices) tasks [[0, 1], [2, 3], [4, 5]] train_loaders [ DataLoader(get_mnist_task(task, True), batch_size64, shuffleTrue) for task in tasks ] test_loaders [ DataLoader(get_mnist_task(task, False), batch_size128, shuffleFalse) for task in tasks ]这个拆分方式属于任务增量学习因为推理时需要传入任务 ID告诉模型应该使用哪一个分类头。后面如果要改成类别增量不能这样给 task ID评估方式会完全不同。3.3 评估指标平均准确率、遗忘率、前向迁移和反向迁移持续学习不能只看“学完最后一个任务后的准确率”还要看整个任务序列上的完整表现。最常用的指标有三个平均准确率Average Accuracy所有已学任务在最终模型上的平均准确率。平均遗忘率Average Forgetting每个任务在刚学完时的准确率减去最终模型在该任务上的准确率差值越大说明遗忘越严重。前向迁移和反向迁移衡量学习一个新任务是否帮助或损害了后续任务、之前任务的表现。具体计算时通常会维护一个准确率矩阵横轴是任务编号纵轴是评估时间点。最后一列对角线是每个任务刚学完时的准确率最后一行是最终模型在所有任务上的准确率。多次运行取平均可以降低随机种子带来的波动。4. 一个最小持续学习示例用 EWC 正则跑通任务序列为了让正则化方法可复现下面给出一个完整的 EWC 最小实现。代码重点不是模型结构多复杂而是展示三件事如何保存旧任务参数、如何计算 Fisher 信息、如何把 EWC 惩罚项加进新任务训练。4.1 模型结构多任务多头分类器因为采用任务增量设置模型使用一个共享的 backbone每个任务对应一个独立的线性分类头。任务 0 训练时创建 head 0任务 1 到来时创建 head 1依此类推。import torch import torch.nn as nn import torch.nn.functional as F import torch.optim as optim class TaskIncrementalMLP(nn.Module): def __init__(self, input_size784, hidden256, num_classes_per_task2): super().__init__() self.hidden hidden self.backbone nn.Sequential( nn.Linear(input_size, hidden), nn.ReLU(), nn.Linear(hidden, hidden), nn.ReLU(), ) self.heads nn.ModuleList() self.num_classes_per_task num_classes_per_task def ensure_head(self, task_id): while len(self.heads) task_id: self.heads.append(nn.Linear(self.hidden, self.num_classes_per_task)) def forward(self, x, task_id): h self.backbone(x.view(x.size(0), -1)) return self.heads[task_id](h)这里的ensure_head是关键。新任务到来时必须先创建对应的分类头否则前向传播会报索引越界。每个任务都使用独立输出头符合任务增量学习的前提推理时能拿到任务 ID。4.2 EWC 的核心公式和代码实现EWC 包含两个核心对象旧任务参数和 Fisher 信息。Fisher 信息通过旧任务数据上梯度的平方近似得到。def compute_fisher(model, loader, task_id, max_samples1000): model.eval() fisher { n: torch.zeros_like(p) for n, p in model.named_parameters() if p.requires_grad } count 0 for x, y in loader: out model(x, task_id) loss F.cross_entropy(out, y) model.zero_grad() loss.backward() for n, p in model.named_parameters(): if p.grad is not None: fisher[n] p.grad.detach() ** 2 count y.size(0) if count max_samples: break for n in fisher: fisher[n] / max(1, count) return fisher这里使用“训练后模型在旧任务样本上的梯度平方”作为 Fisher 近似。严格来说Fisher 信息应该基于对数似然的梯度期望但实际 EWC 实现中很多代码会用交叉熵损失梯度平方代替。这样做在分类问题里是可接受的近似。得到 Fisher 后把旧任务参数保存下来。更新新任务时EWC 惩罚项会限制共享 backbone 和旧任务分类头在重要方向上的移动。def ewc_penalty(model, fisher, old_params, lam1000): penalty 0.0 for n, p in model.named_parameters(): if n in fisher: penalty (fisher[n] * (p - old_params[n]) ** 2).sum() return lam * penalty注意old_params中不会包含后加入的新任务 head 参数所以在遍历当前模型参数时只对n in fisher的参数施加约束即可。这样新任务的分类头可以自由更新而 backbone 和旧任务参数受到 EWC 保护。4.3 训练循环学习任务 1 后继续学习任务 2训练函数本身不复杂。普通增量训练只计算当前任务的交叉熵损失EWC 训练在交叉熵损失基础上加入正则项。def train_task(model, loader, task_id, optimizer, epochs3, fisherNone, old_paramsNone, lam0.0): model.train() for _ in range(epochs): for x, y in loader: out model(x, task_id) loss F.cross_entropy(out, y) if fisher is not None: loss loss ewc_penalty(model, fisher, old_params, lam) optimizer.zero_grad() loss.backward() optimizer.step()评估函数按任务 ID 分别计算准确率。def evaluate(model, task_loaders): model.eval() accs [] with torch.no_grad(): for task_id, loader in enumerate(task_loaders): correct 0 total 0 for x, y in loader: out model(x, task_id) pred out.argmax(dim1) correct (pred y).sum().item() total y.size(0) accs.append(correct / total) return accs主流程如下先训练任务 0训练结束后计算 Fisher 并保存参数之后训练任务 1 和任务 2 时把 Fisher 和旧参数传入train_task。torch.manual_seed(0) model TaskIncrementalMLP() fisher None old_params None for task_id, train_loader in enumerate(train_loaders): model.ensure_head(task_id) optimizer optim.Adam(model.parameters(), lr1e-3) if task_id 0: train_task(model, train_loader, task_id, optimizer, epochs3) fisher compute_fisher(model, train_loader, task_id, max_samples1000) old_params { n: p.detach().clone() for n, p in model.named_parameters() } else: train_task( model, train_loader, task_id, optimizer, epochs3, fisherfisher, old_paramsold_params, lam1000 ) accs evaluate(model, test_loaders[:task_id 1]) print(fafter task {task_id}: {[round(a, 3) for a in accs]})这里为了对比方便任务 0 不使用 EWC任务是确定旧知识基线。任务 1 开始启用 EWC观察旧任务准确率是否比普通增量训练下降得更慢。4.4 验证结果比较普通训练和 EWC 的遗忘程度代码运行结束后会看到类似下面的输出结构。这里只是结构示意不是固定实验结果不同随机种子、不同 epoch、不同lam都会得到不同数值。after task 0: [0.962] after task 1: [0.821, 0.955] after task 2: [0.783, 0.931, 0.947]要对比 EWC 是否有效可以再运行一个lam0的完全普通增量训练。重点看对角线左边的数字任务 0 的准确率在任务 1、任务 2 学完后是否明显降低。如果普通训练下任务 0 准确率明显下滑而 EWC 训练下滑更小说明代码链路是通的EWC 的约束确实在起作用。这个示例还有两个可以改进的方向。第一任务 0 结束后可以也计算 Fisher但不在任务 0 训练时加惩罚第二可以把 Fisher 分任务累积在任务 2 中同时约束任务 0 和任务 1 的重要参数。多任务累积是 EWC 在实际使用中的常见变体。5. 持续学习项目常见的坑和排查路径持续学习实验看起来简单但很多结果失真都藏在数据拆分、评估协议和超参设置里。下面几个坑在持续学习项目里非常常见。5.1 数据划分不一致导致指标失真一个典型的错误是任务 1 的测试集和任务 2 的训练集有重叠。比如任务 1 使用数字 0-4任务 2 也用数字 0-4但任务 2 数据里有大量任务 1 测试样本评估出来的“保持旧知识”能力其实是数据泄漏带来的假象。检查方式很简单在划分任务后打印每个任务的标签集合和样本量确认任务之间的类别边界清晰。如果任务是有业务含义的还要确认时间上没有穿越。比如用上个月数据做训练这个月数据做测试不能把本月数据提前混进训练集。5.2 回放缓冲区被污染或者验证协议不一致使用回放方法时缓冲区里的旧样本必须来自训练集不能来自测试集。另一个常见问题是验证时没有区分任务形态任务增量场景里测试时给 task ID类别增量场景里测试时不给 task ID两种验证方式不能混用。如果实现的是任务增量每个任务有