简介这套基于持续学习的图像分类项目源码面向需要完成机器学习大作业或毕业设计的高校学生也可作为课程设计与初期项目立项的参考。项目针对图像分类中的灾难性遗忘问题以CIFAR100为基础数据集实现了增量任务学习并支持通过命令行参数配置初始类别数、每轮新增类别数、经验重放策略、较小遗忘损失权重等便于扩展对比实验。压缩包共86个文件主体为32个Python源码文件、48个编译后的pyc文件另含4个txt和2个md说明文档整体仅115KB结构清晰。目前已有111人学习下载适合希望在持续学习方向快速上手或二次开发的读者。项目内含主训练/验证脚本、ResNet系列模型、余弦分类器、边缘排序损失与较少遗忘损失模块并附依赖清单与超参说明安装依赖后按示例命令即可运行能帮助理解类增量学习的主流做法。1. 持续学习图像分类做完了图像分类为什么模型还是会“忘本”临近交机器学习大作业很多同学会走那条最熟的路加载一个现成的预训练模型在 10 类或 100 类图像上拿到不错的准确率就算交差。但如果把任务换成“先学 10 类、再学 10 类连续学 10 个任务”你会发现一个让图像分类模型集体翻车的现象新任务学得越好旧任务的准确率掉得越厉害。这个现象叫灾难性遗忘也是持续学习要解决的核心问题。“基于持续学习的图像分类python源码项目说明机器学习大作业.zip”这类题目本质上不是让你改一个分类网络而是让同一个模型在连续到来的数据流里既学得动新类、又忘不了旧类。这个方向适合机器学习期末想冲高分的同学也适合已经跑通常规图像分类、想往真实场景多走一步的开发者代码量不大能讲清楚原理的地方却非常多。2. 持续学习的算法选型EWC、LwF、样本回放的差别在哪儿2.1 普通图像分类的隐藏前提数据分布是静态的常规图像分类之所以很少被人质疑流程有问题是因为它的隐含假设是训练集和测试集来自同一个分布模型只需要在一次性给全的数据里学一次。训练时遇到多少类测试时就是多少类权重从头到尾只训练一次做完就收工。你平时在 Kaggle 或课程实验里跑的分类任务基本都属于这个静态范式。持续学习把这件事改掉了。数据不是一次到齐而是一个任务接一个任务地来。任务 1 给你苹果、香蕉、汽车任务 2 给你猫、狗、飞机任务 3 又出现新类别。每个任务到来时旧任务的原始数据不一定还能访问大作业里通常理解成“旧任务的训练集被收回”。模型必须在不丢掉旧知识的前提下学会辨别新类这是持续学习与普通分类最本质的区别。真实场景里这样的例子很多。一条产线上午拍的零件类别和下午拍的成品类别分属不同生产阶段遥感影像里先做了森林图像分类下个季度又要加裸地和居民区类别。数据不是一次性打包给你的业务却要求同一个模型一直用下去这就是持续学习存在的理由。大作业最常见的评测场景是 Split CIFAR-100把 100 个类别按顺序切成 10 个任务每个任务 10 类。训练完任务 1 再训练任务 2依次类推。评测时不是只看最后一个任务的准确率而是要算模型学完所有任务之后任务 1 的分类能力还剩多少。既然测试集一直在变普通分类里“训练完就测试”的做法就完全不适用了。2.2 灾难性遗忘是怎么发生的共享权重被新任务改坏了假设模型在任务 1 上学好的参数是 θ₁它在任务 1 的特征空间里已经找到了一组好用的组合。任务 2 的梯度下降只会朝着“降低任务 2 损失”的方向更新根本不关心这些更新会不会把任务 1 的损失曲面抬起来。神经网络参数量大、表征高度耦合只要骨干网络共享那么任务 2 的几次大步更新就足够把任务 1 分好的决策边界推到一边去。很多新人第一次复现时把灾难性遗忘误会成“任务 1 的数据被模型删了”。其实数据还在硬盘里网络结构也没变化错的是权重被覆盖。这也是为什么“重放旧数据”看起来最粗暴却一直被用作最强 baseline——它直接在梯度层面挡住了覆盖动作新任务每次反向传播时老样本都在旁边拉一把。持续学习里常说“守旧与学新”的取舍指的就是不知道该把多少梯度预算花在保住旧任务上多少花在学新任务上。EWC、LwF、样本回放这三条路线本质上是三种“怎么分配梯度预算”的策略。它们护旧的手段不同适用场景、代码量和翻车点也差很多。2.3 三条路线怎么选正则化、蒸馏与样本回放持续学习本身不是一套新的图像分类模型而是一套训练策略。核心分类方法仍然可以用 ResNet、ViT 这类常规骨干持续学习算法负责在训练流程上防止权重被覆盖。常见路线有下面三种路线代表方法护旧手段是否依赖旧任务数据大作业代码量参数正则化EWC对重要参数加二次惩罚不需要但要旧任务数据算 Fisher最少约 30 行知识蒸馏LwF用旧模型输出约束新模型不需要中等样本回放iCaRL / Replay保存少量旧样本混入训练需要少量存储较大我一般建议大作业优先选 EWC原理好讲Fisher 信息矩阵的计算在纯 Python 加 PyTorch 环境下只有十几行答辩时能把“为什么重要参数动不得”讲清楚。想再加分就在 EWC 上叠一层 LwF 蒸馏让新任务的 loss 里同时带上旧模型的软目标既在参数层面加了约束又在输出层面兜了底。样本回放路线效果通常最好但要处理“每类存几张”“按什么策略选旧样本”“在哪个阶段重放”这些细节平衡点没调好反而容易翻车。如果大作业题目明确说旧数据已经拿不到那就老老实实用 EWCLwF这两个方法都不依赖旧样本正好符合题目设定。3. 把代码跑起来Split CIFAR-100、Fisher 信息矩阵与 EWC 主流程3.1 先看清项目里每个文件在干什么拿到这类大作业项目别急着运行 train.py先把文件结构过一遍。常见做法是把配置、数据、模型、正则化、训练和评测拆成独立模块project/ ├── config.py # 任务数、每任务类数、训练轮数、EWC 正则系数 lam ├── dataloader.py # 按任务切分 CIFAR-100返回当前任务的 train/test loader ├── model.py # 骨干网络 按任务扩展的输出头 ├── ewc.py # Fisher 信息矩阵计算 EWC 损失 ├── train.py # 主训练循环每个任务结束后保存参数和 Fisher ├── eval.py # 每学完一个任务测一遍所有已见任务记录 acc 矩阵 └── utils.py # 固定随机种子、画混淆矩阵、输出日志其中最容易被忽略的是 eval.py。只记录当前任务测试准确率的项目后期复盘时基本等于黑匣子因为画不出遗忘曲线。如果源码包里的 eval.py 只有单任务测试建议自己补上按任务遍历的逻辑。3.2 数据场景搭建从普通 CIFAR-100 切出 Split CIFAR-100持续学习评测里最常用的数据集是 Split CIFAR-100。它不需要额外下载特殊数据只需要把 torchvision 里的 CIFAR-100 按类别切成多个任务。下面是每个任务生成独立 dataloader 的核心逻辑import torch from torch.utils.data import Subset, DataLoader from torchvision import datasets, transforms # 把 CIFAR-100 按类别顺序切成 num_tasks 个任务返回第 task_id 个任务的数据 def split_cifar100_task(task_id, num_tasks10, batch_size64): classes_per_task 100 // num_tasks start task_id * classes_per_task end start classes_per_task transform transforms.Compose([ transforms.ToTensor(), # CIFAR-100 数据集的全局均值和标准差固定值写死即可 transforms.Normalize((0.5071, 0.4867, 0.4408), (0.2675, 0.2565, 0.2761)), ]) train_ds datasets.CIFAR100(root./data, trainTrue, downloadTrue, transformtransform) test_ds datasets.CIFAR100(root./data, trainFalse, downloadTrue, transformtransform) # 只挑出类标签在 [start, end) 范围内的样本索引 train_idx [i for i, (_, label) in enumerate(train_ds) if start label end] test_idx [i for i, (_, label) in enumerate(test_ds) if start label end] train_loader DataLoader(Subset(train_ds, train_idx), batch_sizebatch_size, shuffleTrue, num_workers2) test_loader DataLoader(Subset(test_ds, test_idx), batch_sizebatch_size, shuffleFalse, num_workers2) return train_loader, test_loader逻辑说明CIFAR-100 的原生标签是 0~99Split CIFAR-100 只是把标签区间切成 10 段。任务 0 用标签 0~9任务 1 用标签 10~19以此类推。train 和 test 两边必须用同样的范围过滤测试集只参与评测绝不参与训练。参数说明task_id 表示当前第几个任务num_tasks 越大每个任务类越少体验曲线越细batch_size 在显存允许时取 64如果某个任务类别数少导致数据量小可以减到 32。注意 Subset 只是保存索引不复制图像数据内存占用很小。3.3 EWC 的核心模块Fisher 信息矩阵怎么算EWC 的护旧逻辑一句话任务 t1 训练时损失里追加一项对旧参数的惩罚惩罚力度由 Fisher 信息矩阵决定。Fisher 告诉我们哪些参数对旧任务重要参数越重要新任务就越不能动它。计算 Fisher 需要在任务 t 的数据上对损失求梯度再取平方均值import torch def compute_fisher(model, dataloader, device): model.eval() criterion torch.nn.CrossEntropyLoss() # 为每一层参数准备一个同形状的 Fisher 张量 fisher {name: torch.zeros_like(p) for name, p in model.named_parameters() if p.requires_grad} for inputs, targets in dataloader: inputs, targets inputs.to(device), targets.to(device) logits model(inputs) loss criterion(logits, targets) model.zero_grad() loss.backward() for name, p in model.named_parameters(): if p.requires_grad and p.grad is not None: fisher[name] p.grad.detach() ** 2 n_samples len(dataloader.dataset) for name in fisher: fisher[name] / n_samples return fisher逻辑说明Fisher 是期望下的梯度平方这里用任务数据上的蒙特卡洛近似把每个样本的梯度的平方累加再取平均。detach() 防止梯度继续回传除以样本数让 Fisher 的数值不随数据集大小变化。参数说明计算 Fisher 时的模型参数状态很关键——必须是在任务 t 训练收敛之后、还没被任务 t1 更新之前计算。顺序反了相当于给错误参数记了重要性EWC 就废了。dataloader 用当前任务的训练集即可不需要额外保存旧数据。3.4 EWC 损失重要参数动得越多罚得越重Fisher 算完之后新旧参数之间每一处差异都要用 Fisher 加权。实现上不复杂def ewc_loss(model, fisher_old, params_old, lam, device): # 对每个参数用 Fisher 加权平方误差作为正则项 reg_loss 0.0 for name, p in model.named_parameters(): if p.requires_grad and name in fisher_old: diff p - params_old[name] reg_loss (fisher_old[name] * diff ** 2).sum() return lam * reg_loss / 2.0逻辑说明lambda 前面的 1/2 是为了求导后系数整洁不影响实质。fisher_old[name] 表示该参数的旧任务重要度重要度越高改动同样大小付出的代价越大params_old[name] 是任务 t 收敛后的参数快照必须深拷贝否则训练器原地更新会把旧参数覆盖掉惩罚项会错误地鼓励新状态和变化后的旧状态一致。常取的 lam 区间是 10~100具体调参见第 4 章。这里只强调一点如果代码里把 lam 设成 0等于整个 EWC 模块没工作要留意配置是否被正确读取。3.5 主训练循环每个任务学完就保存旧参数与 Fisher标准训练循环需要管理三件事当前任务的交叉熵损失、来自历史任务的 EWC 惩罚、以及任务结束后更新 params_old 和 fisher_old# train.py 核心循环模型、优化器、config 按项目实际调整 task_num config.num_tasks # 例如 10 params_old None # 第一个任务没有历史约束 fisher_old None for task_id in range(task_num): train_loader, test_loader split_cifar100_task(task_id, task_num) optimizer torch.optim.SGD(model.parameters(), lr0.01, momentum0.9) for epoch in range(config.epochs_per_task): # 每任务 10~20 轮 for images, labels in train_loader: images, labels images.to(device), labels.to(device) logits model(images) loss torch.nn.functional.cross_entropy(logits, labels) if params_old is not None: # 第 0 个任务不做 EWC loss ewc_loss(model, fisher_old, params_old, config.lam, device) optimizer.zero_grad() loss.backward() optimizer.step() # 任务收敛后保存一份参数快照并计算 Fisher params_old {name: p.detach().clone() for name, p in model.named_parameters() if p.requires_grad} fisher_old compute_fisher(model, train_loader, device) torch.save({model: model.state_dict(), fisher: fisher_old, params_old: params_old}, fcheckpoints/task_{task_id}.pth)逻辑说明第一个任务没有历史可护params_old 为空时 EWC 项不参与计算这是照抄代码时最容易漏掉的边界条件。每个任务训练完再保存 model、fisher、params_old 三样东西checkpoint 是后面的后悔药某组参数跑崩了至少能从中间任务恢复不用从头重跑。参数说明这里用 SGD 起步lr0.01、momentum0.9 是持续学习里比较省心的组合比 Adam 更不容易把旧任务带偏具体原因在第 4 章展开。checkpoint 文件名带 task_id累加实验后比较好定位是哪个任务出了问题。4. 参数调优与实验设计让模型既不遗忘又能学到新类4.1 三个必调超参数EWC 的 λ、蒸馏温度 T、每任务训练轮数EWC 的 λ 是第一优先级。λ0 等价于普通分类旧任务被刷掉是必然λ 给到 10 时已能看到明显约束到 100 时新任务准确率开始下滑个别特征耦合很强的模型λ500 甚至学不进去。常见做法是先跑 λ0 作为下界 baseline再跑 λ10 和 λ100 三组对比观察旧任务准确率是怎么随 λ 回升的。不用追求小数级别的精细值把 0/10/100 三个数量级跑明白答辩时规律就很清楚。如果你在 EWC 上叠了 LwF 蒸馏还需要一个温度 T常见取 2~4。T1 时软目标退化成硬标签蒸馏项等于没有意义T 太高会让所有类别的概率都变得过于平滑旧类别间的差异被一起抹平。蒸馏项的权重 α 一般取 0.5~1.0它与 EWC 损失是并列关系需要两个权重分开控制混在一起很难定位是哪个项在起作用。每任务训练轮数也常被忽略。每任务 5 轮时新任务欠拟合Fisher 算出的重要性不准每任务 20 轮以上新任务过拟合旧任务被遗忘的速度反而更快。我一般先在单任务上定轮数保证“单独学这一个任务”能到 90% 以上再切成连续学习去看整体曲线。Split CIFAR-100 这类数据每任务 10~20 轮是合理区间。4.2 学习率与优化器持续学习的玄学往往出在这里持续学习里影响终点的往往是优化器而不是网络结构。很多同学习惯直接用 Adamlr1e-3前期跑得很顺到任务 3 之后旧任务指标开始抖动。原因是 Adam 会放大每个参数近期的梯度而新任务的梯度天然占主导加上 EWC 正则只加固旧任务的权重、不加固优化器状态“越往后越难锁住旧知识”就成了常事。我的血泪经验是用带动量的 SGDlr 设在 0.005~0.02momentum0.9weight_decay5e-4整体表现更平顺实验也更好复现。学习率调度也要单独设计。常见做法是每个任务内都从固定 lr 开始训练不跨任务保存 scheduler 状态缺点是会丢掉一部分进度。更省心的做法是给每个任务配一个余弦退火或 StepLR让任务内后期小步微调收敛。注意不要在任务切换时突然把学习率降一个数量级——那不是“让旧任务学得更细”而是会干扰正在训练的新任务老任务的准确率不升反降。4.3 评测方式先想清楚多头设置还是单头设置答辩时最容易被问翻车的问题测试这个模型的时候你知道当前样本属于哪一组类别吗如果每个任务有自己的分类头测试时直接把样本送进对应任务头这属于多头宽松评测如果所有任务共用一个分类头测试时完全不提供任务编号这才是更接近真实应用的类增量设置。常见做法是在 config 里加一个 single_head 开关。多头 EWC 在 10 个任务上的最终平均准确率通常比单头高 10 个百分点以上这不是模型变厉害了而是评测标准放松了。写大作业报告时必须注明评测是哪种否则你报告里的 70% 和学长报告里的 55% 根本不可比老师一眼就能看出来。实验表里两条腿一起跑先报多头再说单头两个数字都有才是完整数据。4.4 实验记录表把每组参数和指标留档跑持续学习实验很容易跑乱同一份代码改一个 λ结果记录就找不到了。建议每组实验建一个独立目录名字带日期和关键参数同时用一张表记录所有配置与结果实验编号骨干网络方法lamT优化器lr每任务轮数多头ACC单头ACC遗忘率FGTE001ResNet18EWC0—SGD0.011563.2048.1032.40E002ResNet18EWC10—SGD0.011570.4055.8018.60E003ResNet18EWCLwF103SGD0.011572.1058.3015.20每次实验跑完把 acc 矩阵保存成 npy 文件和实验表放同一目录。这样后期画遗忘率曲线时不用再重新跑一遍也能在复盘时定位是哪一组参数导致指标变化。固定随机种子这一项也要写进表里不固定种子的实验数据互相之间没有可比性。5. 避坑指南持续学习图像分类最容易踩的 5 个坑5.1 数据切分翻车测试集混进了训练任务现象训练 10 个任务后每个任务单独测准确率都很高但把混在一起的测试集拿去画混淆矩阵整个矩阵乱成一团明显不对。原因最常见是把 torchvision 的 CIFAR-100 训练集按任务切了测试集却没按同样的类区间过滤。eval 时拿全量测试集喂给只见过部分类的模型模型对没见过类的输出完全是乱的指标当然没有意义。解决train 和 test 必须按完全相同的类区间过滤也就是第 3.2 节代码里那样两组 idx 都做。再检查一处确认 eval 时用的是当前任务的 test_loader而不是复用 train_loader。5.2 Fisher 信息矩阵只存了最后一个任务现象EWC 代码看起来没问题跑完 10 个任务后任务 1 的准确率比 λ0 的 baseline 还低正则项像个摆设。原因每训练完一个任务就把 fisher_old 整体覆盖成新任务的 Fisher旧任务的重要性信息全丢了。任务 3 训练时只对任务 2 做了约束任务 1 完全没有保护。解决Fisher 要跨任务累加或者保留每个任务的 Fisher 求和。简单做法是fisher[name] fisher[name] fisher_new[name]而不是直接赋值。严格来说这是 EWC 的累加版实现上只差一行效果却差很多。如果嫌旧任务 Fisher 拖累新任务学习可以乘一个衰减系数但大作业阶段直接累加最稳。5.3 BN 层在训练和评测之间统计量漂移现象训练时 loss 正常下降每任务单独测也在进步但把模型切到 eval 模式后某些旧任务准确率突然掉了一截换 batch_size 之后掉得还不一样。原因BN 层维护的是 running mean 和 running var任务切换后这些统计量继续被新任务更新。当任务间图像风格差异大、或者 batch 比较小时统计量会向新任务偏移旧任务在 eval 时用的分布已经不准了。解决batch_size 尽量不低于 64在训练新任务时对已经收敛好的旧任务 BN 参数做冻结只更新新任务有关的层。更省心的做法是把骨干网络的 BN 替换成 GroupNorm少一个统计量就少一个坑小 batch 下效果更稳定。这个替换在代码里可以在 model.py 里用一行循环完成副作用是可能要把学习率略微调低一点。5.4 随机种子没固定实验对不齐现象同一份代码、同一组参数跑两次结果差 5 个百分点以上报告没法写老师复现也对不上。原因dataloader 的 shuffle、dropout、卷积算子选择都带随机性。没有固定随机种子持续学习里任何一次随机波动都会沿着任务链条往后放大越到后面的任务差异越明显。解决在 utils.py 里统一固定 random、numpy、torch、cudnn 的种子def set_seed(seed42): import random import numpy as np import torch random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) # 关闭自动选择算法保证卷积算子一致 torch.backends.cudnn.deterministic True torch.backends.cudnn.benchmark False参数说明deterministicTrue 会把部分卷积算子换成确定版本训练速度会慢 8%~20%换来的是逐位可复现大作业阶段值得。benchmarkFalse 是为了避免输入尺寸变化时重新选算法带来的随机性。这段代码必须在所有数据加载和模型初始化之前调用放在 import 之后的 main 开头即可。5.5 日志只记 total loss最后画不出遗忘曲线现象一个又一个任务跑完想复盘时发现手里的记录只有每轮的 total loss 和最后一个任务的准确率遗忘曲线没数据画论文里缺关键图。原因eval.py 只测了当前任务或者把评估代码放在了训练循环外面没在任务切换时对所有已见任务做一遍评测。这等于把实验的中间状态全部丢掉后期只能重跑。解决维护一个 acc_matrix每学完一个任务就遍历所有已见任务计算准确率并同时保存成 npy 和文本文件acc_matrix np.zeros((num_tasks, num_tasks)) for task_id in range(num_tasks): # 训练任务 task_id 的代码略... for eval_task in range(task_id 1): acc_matrix[task_id, eval_task] compute_acc(model, eval_task) np.save(logs/acc_matrix.npy, acc_matrix)逻辑说明acc_matrix[i, j] 表示学完任务 i 后测任务 j 的准确率这是持续学习论文最标准的报告格式。保存成 npy 后第 6 章的遗忘率和平均准确率直接从这个矩阵算不用重新加载模型。6. 验证与进阶用遗忘率曲线证明模型真的在“持续学习”6.1 平均准确率与遗忘率的计算持续学习实验报数就两个核心指标学完所有任务后已见任务的平均准确率 ACC以及旧任务比刚学完时掉了多少的平均值叫遗忘率 FGT。计算公式如下import numpy as np def compute_metrics(acc_matrix): # acc_matrix[i, j]: 学完任务 i 后测试任务 j 的平均准确率 T acc_matrix.shape[0] # 最后一个任务学完后所有已见任务的平均准确率 acc acc_matrix[T - 1, :T].mean() # 每个任务在刚学完时通常达到最高点用它减去最终准确率求遗忘 best np.max(acc_matrix[:T, :T], axis0) final acc_matrix[T - 1, :T] forget (best - final)[:T - 1].mean() return acc, forget逻辑说明有两个细节要留意。final 是 acc_matrix 最后一行所有任务的值best 是对角邻近区间的局部峰值——一个任务刚学完时的准确率实际上是前一刻还没有后续任务来破坏它的状态所以峰值通常出现在它本身学完那一行或者紧接着的一两行。用这个公式能算出每个旧任务被遗忘了多少再取平均就是 FGT。注意最后一行代码的[:T-1]最后一个任务没有后续任务来“遗忘”它如果公式里不加这个切片遗忘率会被最后一个任务的噪声抬高评审抓住这一点就麻烦了。6.2 按任务分块的混淆矩阵能看出什么除了两个数值指标按任务分块的混淆矩阵更能说明问题。把模型在所有已见类别上的预测画成一个大混淆矩阵然后沿任务边界画出分块线。如果任务 1 的样本大量错分到任务 2 的类别区域说明任务 2 的训练把底层特征带偏了光看平均准确率发现不了这种方向性偏移。分块矩阵的另一个用途是检验数据切分是否正确正常时对角块最亮相邻任务间的浅色块分布均匀。如果某个任务对应的整行都暗基本就是那组任务的 BN 统计量出了问题或 Fisher 累加时把它漏掉了。6.3 往真实项目走的三个进阶方向大作业验收完之后还想让这个方向更有分量可以按下面三个方向各走一步。第一换预训练骨干。用 ImageNet 上预训练的 ResNet18 替换随机初始化模型旧任务遗忘率会大幅下降这是常识性结论但报告里必须补一组随机初始化对照否则老师会质疑你的提升来自方法还是来自预训练。第二给 EWC 加少量回放。每类保存十几张真实旧样本在训练新任务时混入 batch正则加回放的组合非常接近真实系统里“少量旧数据还能拿到”的设定。第三把多头评测换成完整的类增量评测即测试时不给任务编号这一条能让报告的说服力上一个台阶。我自己跑这类实验养成的习惯是每组实验一个带日期的目录checkpoint、acc_matrix、实验记录表三样东西永远放一起跑完立刻生成 npy 和一张简洁的曲线图绝不留到第二天补。这个习惯救过我很多次因为持续学习实验一跑就是几十个小时忘了保存中间状态重新跑的代价不是时间而是对“到底哪组参数有效”的判断会失真。图像分类入门容易持续学习要做出可信的结果需要细心希望这篇笔记能帮你少走一段弯路也希望你做完之后能把这套流程用到下一个更复杂的任务上。本文还有配套的精品资源点击获取