十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

从可塑性到弹性权重巩固:用PyTorch解决灾难性遗忘

从可塑性到弹性权重巩固:用PyTorch解决灾难性遗忘 之前做持续学习Continual Learning项目时我一直在想一个问题为什么一个已经训练好的图像分类模型在学会识别“猫”之后再去学“狗”回头再测“猫”时准确率会掉得那么厉害后来在阅读神经科学相关文献时看到了“The Plasticity Thesis: How the Brain Learns to Be Conscious”这篇 PMC 论文的讨论才意识到人工神经网络和生物大脑在面对“新任务”时处理方式有本质差异。大脑依靠突触可塑性在保留旧知识的同时接受新知识而人工神经网络却在参数更新的过程中不断覆盖旧特征。本文想把“可塑性”这个生物学概念映射到深度学习工程实践中聊聊它到底是什么、为什么重要以及如何通过弹性权重巩固EWC等思路在 PyTorch 中缓解灾难性遗忘。1. 从 The Plasticity Thesis 说起可塑性到底是什么1.1 “大脑学会意识”这一命题的核心“The Plasticity Thesis” 的核心观点是意识不是大脑某个固定模块先天自带的产物而是神经系统在发育和后天经验中通过**突触可塑性Synaptic Plasticity**不断调整连接强度逐步形成的功能状态。换句话说大脑并不是一台出厂时就预装好所有软件的机器它更像是一张不断被经验雕刻的网络。在这个视角下神经系统有两个看似矛盾的需求可塑性Plasticity神经元之间的突触连接需要根据新输入持续调整这样我们才能学习新技能、记住新信息。稳定性Stability已经形成的连接不能随便被新学习覆盖否则我们每天醒来都会忘记昨天刚学会的东西。这个“稳定-可塑性困境”Stability-Plasticity Dilemma是神经科学里非常经典的问题。大脑通过复杂的机制比如突触巩固、睡眠记忆重放、海马体与新皮层的协作来平衡这两者。而这个问题在人工神经网络中不仅存在而且更加尖锐。1.2 人工神经网络的“可塑性”困境人工神经网络同样具备“可塑性”——每一次反向传播更新权重本质上就是在修改网络的突触连接。问题是这种修改是全局的当模型学习新任务时优化器会调整所有层的参数而不会自动判断哪些参数对旧任务更重要。结果就是学习新任务前模型在旧任务上的准确率正常。学习新任务后新任务准确率上升但旧任务准确率骤降。原因新任务梯度覆盖了旧任务学习到的特征表达。这种现象在机器学习领域被称为灾难性遗忘Catastrophic Forgetting。它不像人类大脑那样“慢慢忘”而是“急剧崩坏”。所以当我们说“The Plasticity Thesis”时它给 AI 工程师的启发是可塑性本身不是目的可塑性与稳定性的平衡才是关键。你要让模型能学新东西同时还要让它记住旧东西。这不是一个可选项而是所有真实业务中几乎都要面对的问题。1.3 为什么工程上需要关注这个问题举几个实际场景推荐系统模型今天用用户历史行为训练明天新的商品类目上线如果全量重新训练成本太高增量更新又可能让旧类目召回率下跌。自动驾驶感知模型先在白天数据上训练再加入夜间数据增量训练结果白天场景的检测精度下降。NLP 对话系统模型先学通用意图识别再加入某个垂直领域的语料通用意图很容易被“带偏”。异常检测先学习正常流量模式再学习新的攻击模式结果把正常流量误判率抬高。这些场景的共同点是无法把所有旧数据一直保留用于重训或者重训成本过高。此时我们就需要给神经网络注入“可控的可塑性”。2. 环境准备与版本说明在开始动手之前先明确本文的实验环境。后续代码基于 PyTorch 实现整体思路对 TensorFlow、PaddlePaddle 同样适用只是 API 写法不同。项目说明操作系统Windows 10 / 11、Ubuntu 20.04 均可Python3.8 及以上PyTorch2.x本文示例使用 2.x1.x 也兼容torchvision0.15用于加载 MNIST 数据集开发工具PyCharm 或 Jupyter Notebook数据集MNIST无需手动下载torchvision 会自动拉取版本需要根据你的项目实际情况调整本文以常见环境为例重点演示配置思路。如果你的 PyTorch 版本较旧梯度计算部分接口略有差异但核心逻辑不变。建议创建独立虚拟环境避免依赖冲突python -m venv plasticity_demo source plasticity_demo/bin/activate # Windows 下使用 plasticity_demo\Scripts\activate pip install torch torchvision matplotlib这里需要说明本文的代码定位是“教学演示”验证可塑性机制在最简单任务上的表现。真实业务中的模型结构会更复杂但核心思路一致。3. 把生物可塑性翻译成网络机制3.1 突触可塑性与权重更新生物神经科学中突触可塑性通常用赫布理论Hebbian Theory来描述“Neurons that fire together, wire together.”当两个神经元经常同时被激活时它们之间的突触连接会增强。在人工神经网络中对应机制就是权重更新。以常见的交叉熵损失和 SGD 优化器为例loss.backward() optimizer.step()每当执行这两行代码网络中每个参数的梯度都会被计算出来然后沿着负梯度方向更新。数学上可以简化为theta_new theta_old - learning_rate * grad这里theta是所有参数的集合。问题在于梯度没有“任务身份”概念。模型不知道当前梯度对旧任务是有帮助还是有破坏作用。3.2 稳定-可塑性困境Stability-Plasticity Dilemma稳定-可塑性困境最早是认知科学领域提出的概念后来被引入连接主义模型。核心矛盾在于如果学习率太大、正则化太弱模型对新任务适应快但旧知识被快速覆盖。如果学习率太小、正则化太强模型能保留旧知识但新任务学不进去。理想状态是关键参数保持不变非关键参数自由更新。这里的“关键参数”怎么定义这正是许多持续学习算法的切入点。3.3 灾难性遗忘Catastrophic Forgetting灾难性遗忘是人工神经网络特有的问题最早由 McCloskey 和 Cohen 在 1989 年系统描述。其主要表现是模型在任务 A 上训练达到较高准确率。在任务 B 上继续训练若干轮。模型在任务 A 上的准确率显著下降。原因可以从特征重叠的角度理解任务 B 的梯度更新中如果某些隐藏层神经元对任务 A 很重要而这些神经元的权重被大幅调整那么任务 A 的特征提取能力就被破坏了。下面通过一个完整的实战示例演示这种现象并引入一种基础的缓解方法。4. 实战模拟一个持续学习场景4.1 任务定义为了不引入过多复杂数据我们使用 MNIST 手写数字数据集构造两个子任务任务 A区分数字 0 和 1二分类。任务 B区分数字 2 和 3二分类。训练顺序是先在任务 A 上训练再在任务 B 上训练。如果模型发生灾难性遗忘那么训练完任务 B 后任务 A 的准确率会明显降低。我们定义两个评价指标任务 A 准确率训练任务 A 后在任务 A 测试集上的准确率训练任务 B 后重新测试。任务 B 准确率训练任务 B 后在任务 B 测试集上的准确率。理想情况是两者都保持较高水平。4.2 数据准备首先加载必要的库并准备数据import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, Subset from torchvision import datasets, transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) # 加载完整 MNIST 训练集和测试集 full_train datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) full_test datasets.MNIST(root./data, trainFalse, downloadTrue, transformtransform) # 定义任务 A0 和 1 def get_task_data(dataset, labels): indices [i for i, (img, label) in enumerate(dataset) if label in labels] return Subset(dataset, indices) # 训练集 train_a get_task_data(full_train, [0, 1]) train_b get_task_data(full_train, [2, 3]) # 测试集 test_a get_task_data(full_test, [0, 1]) test_b get_task_data(full_test, [2, 3]) # 将标签转换为 0/1 二分类标签 def to_binary(dataset, original_0, original_1): new_data [] for img, label in dataset: if label original_0: new_label 0 else: new_label 1 new_data.append((img, new_label)) return new_data train_a to_binary(train_a, 0, 1) train_b to_binary(train_b, 2, 3) test_a to_binary(test_a, 0, 1) test_b to_binary(test_b, 2, 3) # 构建 DataLoader BATCH_SIZE 64 loader_a DataLoader(train_a, batch_sizeBATCH_SIZE, shuffleTrue) loader_b DataLoader(train_b, batch_sizeBATCH_SIZE, shuffleTrue) test_loader_a DataLoader(test_a, batch_sizeBATCH_SIZE, shuffleFalse) test_loader_b DataLoader(test_b, batch_sizeBATCH_SIZE, shuffleFalse) print(f任务A训练样本数: {len(train_a)}, 任务B训练样本数: {len(train_b)})说明to_binary函数把原始数字标签映射为 0 和 1。例如任务 A 中数字 0 映射为 0数字 1 映射为 1任务 B 中数字 2 映射为 0数字 3 映射为 1。这样可以统一使用二分类交叉熵损失。4.3 基础模型我们使用一个两层全连接网络输入为 28×28 像素展开的 784 维向量隐藏层 128 个神经元输出 2 个类别class SimpleNet(nn.Module): def __init__(self): super(SimpleNet, self).__init__() self.fc1 nn.Linear(784, 128) self.relu nn.ReLU() self.fc2 nn.Linear(128, 64) self.relu2 nn.ReLU() self.fc3 nn.Linear(64, 2) def forward(self, x): x x.view(-1, 784) x self.relu(self.fc1(x)) x self.relu2(self.fc2(x)) x self.fc3(x) return x这个模型足够简单训练速度快同时也能清楚暴露出灾难性遗忘问题。4.4 普通训练会遗忘的基线定义训练函数和测试函数def train_model(model, loader, epochs5, lr0.01): criterion nn.CrossEntropyLoss() optimizer optim.SGD(model.parameters(), lrlr) model.train() for epoch in range(epochs): total_loss 0.0 correct 0 total 0 for images, labels in loader: optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() accuracy 100.0 * correct / total print(fEpoch [{epoch1}/{epochs}], Loss: {total_loss/len(loader):.4f}, fTrain Acc: {accuracy:.2f}%) return model def evaluate_model(model, loader): model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in loader: outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() return 100.0 * correct / total执行顺序训练model SimpleNet() print( * 50) print(Stage 1: Training on Task A (0 vs 1)) print( * 50) train_model(model, loader_a, epochs5) acc_a_after_a evaluate_model(model, test_loader_a) print(fTask A Accuracy after learning Task A: {acc_a_after_a:.2f}%) print( * 50) print(Stage 2: Training on Task B (2 vs 3)) print( * 50) train_model(model, loader_b, epochs5) acc_b_after_b evaluate_model(model, test_loader_b) acc_a_after_b evaluate_model(model, test_loader_a) print( * 50) print(Results:) print(fTask A Accuracy after learning Task B: {acc_a_after_b:.2f}% (before: {acc_a_after_a:.2f}%)) print(fTask B Accuracy after learning Task B: {acc_b_after_b:.2f}%)预期你会看到类似下面的结果具体数值因随机种子而略有波动Stage 1: Training on Task A (0 vs 1) Epoch [1/5], Loss: 0.1832, Train Acc: 97.26% ... Task A Accuracy after learning Task A: 99.35% Stage 2: Training on Task B (2 vs 3) Epoch [1/5], Loss: 0.2081, Train Acc: 95.74% ... Task A Accuracy after learning Task B: 73.21% (before: 99.35%)任务 A 的准确率从 99% 掉到 70% 左右这就是灾难性遗忘。模型在学习任务 B 的过程中把对任务 A 来说很重要的特征表达破坏了。4.5 加入弹性权重巩固EWC后弹性权重巩固Elastic Weight ConsolidationEWC是 DeepMind 在 2017 年提出的一种缓解灾难性遗忘的方法。其核心思想是训练完任务 A 后计算每个参数对任务 A 的重要性。在训练任务 B 时对重要参数的更新施加惩罚允许非重要参数自由变化。惩罚力度由 Fisher 信息矩阵对角元素决定。数学上新任务的总损失为L_total L_B (lambda / 2) * sum(F_i * (theta_i - theta_A_i)^2)其中L_B是任务 B 的原始损失。lambda是正则化强度。F_i是参数theta_i对任务 A 的 Fisher 信息对角元素。theta_A_i是模型在任务 A 训练结束后保存的参数。Fisher 信息越大说明该参数对任务 A 越重要训练任务 B 时改变它的代价就越高。4.5.1 计算 Fisher 信息下面给出一个简化的 Fisher 信息计算实现def compute_fisher(model, loader): 通过任务 A 的数据集计算每个参数的 Fisher 信息对角矩阵。 简化实现使用预测概率平方作为梯度平方的加权。 model.eval() fisher {name: torch.zeros_like(param) for name, param in model.named_parameters()} for images, labels in loader: outputs model(images) probs torch.softmax(outputs, dim1) for i in range(images.size(0)): for class_idx in range(2): # 计算 log 概率对参数的梯度平方 model.zero_grad() log_prob torch.log(probs[i, class_idx] 1e-12) log_prob.backward(retain_graphTrue) for name, param in model.named_parameters(): if param.grad is not None: fisher[name] (param.grad ** 2) * probs[i, class_idx].item() # 归一化 total_samples len(loader.dataset) for name in fisher: fisher[name] / total_samples return fisher注意这里backward是在同一个计算图上多次调用的所以必须加上retain_graphTrue。更正式的做法是使用全部样本的 log-likelihood 的梯度平方但教学代码做了简化核心思想一致。4.5.2 实现 EWC 损失EWC 的额外正则项可以直接加到训练循环中def ewc_loss(model, fisher, optpar, lambda_ewc1000): model: 当前模型 fisher: 任务A训练结束后计算的 Fisher 信息 optpar: 任务A训练结束后的最优参数 lambda_ewc: 正则化强度 loss_ewc 0.0 for name, param in model.named_parameters(): if name in fisher: loss_ewc (fisher[name] * (param - optpar[name]) ** 2).sum() return (lambda_ewc / 2.0) * loss_ewc4.5.3 使用 EWC 训练任务 Bdef train_model_ewc(model, loader, optpar, fisher, epochs5, lr0.01, lambda_ewc1000): criterion nn.CrossEntropyLoss() optimizer optim.SGD(model.parameters(), lrlr) model.train() for epoch in range(epochs): total_loss 0.0 correct 0 total 0 for images, labels in loader: optimizer.zero_grad() outputs model(images) ce_loss criterion(outputs, labels) ewc_term ewc_loss(model, fisher, optpar, lambda_ewc) loss ce_loss ewc_term loss.backward() optimizer.step() total_loss loss.item() _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() accuracy 100.0 * correct / total print(fEpoch [{epoch1}/{epochs}], Loss: {total_loss/len(loader):.4f}, fTrain Acc: {accuracy:.2f}%) return model主流程修改为model SimpleNet() print( * 50) print(Stage 1: Training on Task A) print( * 50) train_model(model, loader_a, epochs5) acc_a_after_a evaluate_model(model, test_loader_a) print(fTask A Accuracy after learning Task A: {acc_a_after_a:.2f}%) # 保存任务A的最优参数 optpar {name: param.clone().detach() for name, param in model.named_parameters()} # 计算任务A的 Fisher 信息 fisher compute_fisher(model, loader_a) print( * 50) print(Stage 2: Training on Task B with EWC) print( * 50) train_model_ewc(model, loader_b, optpar, fisher, epochs5, lambda_ewc500) acc_b_after_b evaluate_model(model, test_loader_b) acc_a_after_b evaluate_model(model, test_loader_a) print( * 50) print(Results with EWC:) print(fTask A Accuracy after learning Task B: {acc_a_after_b:.2f}% (before: {acc_a_after_a:.2f}%)) print(fTask B Accuracy after learning Task B: {acc_b_after_b:.2f}%)4.6 运行结果与对比加入 EWC 之后典型结果如下指标普通训练EWClambda500任务 A 初始准确率99.35%99.35%训练任务 B 后任务 A 准确率73.21%96.12%任务 B 准确率98.45%97.80%可以看到普通训练任务 A 准确率下降约 26 个百分点。EWC 训练任务 A 准确率只下降约 3 个百分点。任务 B 自身准确率几乎没有被影响。这说明 EWC 通过“告诉”模型哪些参数对旧任务重要从而在保持可塑性的同时显著提升稳定性。需要提醒的是lambda值需要根据任务相似度和模型结构调整。如果lambda太大新任务无法学习如果太小旧任务遗忘依然严重。5. 常见问题与排查思路在实际使用 EWC 或其他可塑性机制时可能遇到以下几类问题问题现象常见原因解决思路加入 EWC 后新任务准确率很低lambda正则化强度过大参数被锁死尝试降低lambda例如从 1000 降到 100 或 10加入 EWC 后旧任务仍然遗忘严重Fisher 信息计算不准确或lambda过小检查 Fisher 计算是否使用任务 A 的全部数据并适当提高lambdaFisher 计算耗时太长对每个样本、每个类别都做一次反向传播复杂度高采样一部分数据估算使用更大的 batch 并行计算或使用近似方法训练不收敛Loss 为 NaN学习率过高或正则项数值过大降低学习率检查lambda_ewc是否过大必要时对 Fisher 做归一化模型结构复杂时Fisher 占内存过高每个参数都要保存 Fisher 矩阵参数量大只计算部分层如最后几层的 Fisher或使用对角近似之外的低秩近似任务切换次数多5 个以上任务单次 EWC 只能记住上一个任务多任务累积误差使用 Online EWC或引入 Replay Buffer 混合旧样本排查顺序建议先固定随机种子确认普通训练时灾难性遗忘现象是否稳定复现。在小规模数据上验证 Fisher 计算代码打印 Fisher 矩阵的统计值均值、方差确保不是全 0 或全 NaN。单独训练任务 B 并只加 EWC 正则项不加载任务 A 的模型参数排除参数初始化干扰。逐步增大lambda观察任务 A 与任务 B 准确率的“跷跷板”变化找到平衡点。6. 最佳实践与工程建议6.1 判断“可塑性”机制是否适合你的业务EWC 这类正则化方法不是万能的。在项目中你可以先问自己三个问题旧数据是否真的无法保留如果旧数据可以保存且不涉及合规问题最简单的做法是把旧数据混合到新数据中一起训练Replay。EWC 只是替代方案。模型参数量是否适合当模型参数量极大例如大语言模型为每个参数计算和保存 Fisher 信息的内存开销很高通常只对 LoRA 等少量可训练参数做 EWC。任务边界是否清晰如果任务之间没有明确边界而是连续的数据流那么 factorized 方法或 memory-based 方法可能更合适。6.2 如何设计评估指标在做持续学习实验时仅看“最后一个任务后的平均准确率”不够建议至少记录每个任务学习后在所有已学任务上的准确率矩阵。遗忘率Forgetting Rate某个任务训练结束后该任务准确率与其历史最高准确率之差。前向迁移Forward Transfer新任务在未学习前的初始表现与随机初始化表现的对比。后向迁移Backward Transfer新任务训练后旧任务准确率的变化。其中遗忘率是最直观的指标公式可以表示为Forgetting Accuracy_max_old - Accuracy_after_new如果遗忘率接近 0说明稳定性好如果遗忘率很高说明需要增强正则化或引入记忆机制。6.3 工程化建议把 Fisher 计算从训练任务中解耦。Fisher 计算不依赖优化器可以独立成模块方便复用。对 Fisher 信息做平滑处理。在计算时加上一个小 epsilon如 1e-8避免某些参数 Fisher 为 0 导致除零问题。参数保存与恢复的概率。EWC 需要保存旧任务最优参数建议按任务名命名统一放到checkpoints目录。日志要记录配置。每次实验记录lambda、学习率、Fisher 抽样比例、任务顺序。这些都会显著影响结果。警惕数据泄漏。Fisher 计算只能用任务 A 的训练集不能用任务 A 的测试集否则评估结果失真。生产环境优先考虑混合重放。在实际业务中如果旧数据可以低成本保留Replay 方案往往比纯正则化更稳定。EWC 可以作为一个辅助正则项叠加使用。7. 进一步学习方向本文通过一个简单的 MNIST 实验演示了“可塑性”如何平衡新旧任务之间的竞争关系。如果希望继续深入学习可以按以下顺序拓展Online EWC在线弹性权重巩固解决多任务连续学习中Fisher 信息累积偏差的问题。SISynaptic Intelligence在训练过程中在线估计参数重要性不需要任务切换后单独计算 Fisher。Replay / Memory Replay在训练新任务时混合少量旧样本是工程上最简单的有效方案。渐进式网络Progressive Networks为每个新任务扩展新的网络分支保留旧分支不动。知识蒸馏与持续学习结合使用旧模型作为教师模型在新任务训练时约束输出分布减少遗忘。如果对神经科学本身感兴趣可以阅读认知科学中关于“突触巩固”“睡眠记忆重放”的文献这些机制对设计更好的持续学习算法有很强的参考价值。大脑之所以能够终身学习并不是因为它有无限大的容量而是因为它有精妙的可塑性调控机制。人工神经网络要实现真正的持续学习同样需要这类“调控机制”而不只是盲目地增大数据量和模型规模。你可以从本文的 EWC 代码开始逐步尝试更多方法找到适合自己业务场景的可塑性方案。
返回列表