LUCIR增量学习实战:基于余弦归一化与特征约束缓解灾难性遗忘
2026/9/24 18:48:56 网站建设 项目流程

简介:面向图像分类持续学习任务的Python源码与项目说明,适合计算机、人工智能、数据科学等专业的在校生用于机器学习课程大作业、毕业设计或初期项目演示,也便于入门者理解增量学习中的灾难性遗忘问题。项目基于CIFAR100数据集设计多任务增量训练流程,提供经验重放、较小遗忘损失、边缘排序损失等可选策略,并通过--dataset、--start、--increment、--rehearsal、--selection等参数灵活控制实验配置,方便复现与对比。资源为ZIP压缩包,共86个文件,以32个Python脚本和48个pyc编译文件为主,另有4个文本文件和2个Markdown说明文档,整体仅115KB,轻量精简;Python代码覆盖模型构建、训练、验证、样本选择等模块,Markdown文档则给出项目说明与运行指引。已有111人学习下载,适合需要快速获取可运行持续学习基线代码的读者。解压后按英文路径运行,安装依赖并执行python main.py即可开始实验,也可基于现有模块扩展自定义数据集或损失函数,用于进一步研究与实践。

1. 先让增量学习落地:这份 LUCIR 源码解决的是什么

图像分类做到 90% 以上准确率在今天不算新闻,但让模型在一个任务上训练完,再去学新任务时不把旧知识忘光,这事的难度就完全不一样了。持续学习(Continual Learning)要解决的核心就是灾难性遗忘——神经网络在新数据上做梯度更新时,旧任务的特征表征会被逐渐覆盖。这份基于 LUCIR(Large Margin Cosine Loss with Less-Forget Constraint)的 Python 源码,就是一套在 CIFAR100 上完整实现增量图像分类的工程,适合机器学习大作业、课程设计或毕设起步。它不是你常见的单次训练脚本,而是把「旧样本回放 + 余弦归一化分类器 + 两个约束损失」三件事同时做进去的完整管线,跑通了它能直接作为增量学习方向的实验基座。适合对 PyTorch 有一定基础、想从理论走向可复现实验的在校学生和工程师。

2. 原理先立住:LUCIR 三件套为什么能缓解遗忘

2.1 增量学习的核心矛盾:稳定性与可塑性

传统训练里,模型一次性看到所有类别,梯度更新没有新旧之分。但增量学习的每一轮只给当前任务的类别数据,比如 CIFAR100 先学 50 类,再学 10 类、10 类地往上加。这里有个天然矛盾:要让模型学新类,就得更新权重;但权重一更新,旧类的特征分布就被扰动。

LUCIR 的做法不是简单地把新旧数据混在一起重训,而是从三个层面同时下手:特征表示层面用余弦归一化替代内积分类,让新旧类在超球面上各自占据角度空间;损失函数层面用 Less-Forget 约束锁住旧类特征不漂移;样本管理层面用 Herding 算法挑最有代表性的旧样本做回放。这三件事互相配合,缺一个效果都会明显打折。

为什么用余弦分类器而不是常规的 Linear 层?因为常规内积分类器的权重模长和角度耦合在一起,新类学多了会挤压旧类的决策边界。而把特征和权重都归一化到单位超球面上,分类决策只由夹角决定,新类加入时只是在新角度区域划边界,对旧类区域的干扰要小得多。这份代码里的CosineClassifier.py做的就是这件事。

2.2 Less-Forget 约束:给旧特征加锚点

纯靠余弦分类器还不够,因为特征提取器(backbone)的底层参数仍然会被新任务带动。Less-Forget 的核心思路是:在训练新任务时,把当前模型对旧样本提取的特征,约束在旧模型提取的特征附近。

loss/less_forget.py里的简化逻辑:

def less_forget_loss(current_feats, old_feats, old_classes_mask): # current_feats: 当前模型对旧样本提取的特征 # old_feats: 旧模型冻结后对同一批样本提取的特征 # old_classes_mask: 只对旧类别对应的logits计算约束 diff = current_feats - old_feats.detach() l2_loss = torch.norm(diff, p=2, dim=1) masked_loss = (l2_loss * old_classes_mask).sum() / old_classes_mask.sum() return masked_loss

old_feats.detach()是关键,旧模型的特征只作为监督信号,不参与梯度回传,否则旧模型本身也在变,锚点就失效了。old_classes_mask保证这个损失只约束旧类别对应的特征维度,新类别还在自由学习。lambda_base参数控制这个损失的权重,项目默认配置里一般取 5 到 10,经验上看太小锁不住特征,太大会让新类学不动。

2.3 Margin Ranking Loss:把新旧类的边界推开

Less-Forget 解决的是旧类内部特征的稳定性,但新旧类之间还有另一个问题——新类刚加入时,分类器容易被旧类的大 logits 压制。LUCIR 用 margin ranking loss 拉开新旧正负样本对的距离。

损失本身的逻辑是:从旧样本池里选真正的旧类样本作为正样本,从当前批次的新类里选负样本,要求正负样本之间的余弦相似度差大于一个 margin。这样旧类特征不会贴到新类边界上,分类边界更干净。

def margin_ranking_loss(pos_sim, neg_sim, margin=0.5): # pos_sim: 正样本对相似度 # neg_sim: 负样本对相似度 # 期望 pos_sim - neg_sim > margin,否则产生loss loss = torch.relu(neg_sim - pos_sim + margin) return loss.mean()

margin一般取 0.5 左右,太大会让训练不稳定,太小边界区分度不够。这份代码里loss/margin_lucir.py把 margin ranking loss 和交叉熵整合在一起,训练时两个损失同时回传。

提示:如果你只看论文跑实验,建议先按默认参数完整跑一遍 CIFAR100 的 50+5*10 划分,再分别关闭--less_forg--ranking做消融,能直观看到每个组件对最终准确率的贡献。

3. 把工程跑起来:目录结构、环境与第一次训练

3.1 解压后的目录全景

下载解压后你会看到两层结构,根目录是主工程,source_code_all_bk是备份副本。主目录里main.py是入口,train.pyvalidate.py负责训练验证,models/下是网络结构,loss/下是两个约束损失实现,utils/下是样本选择和数据管理工具。

注意:项目名和路径不要用中文,否则容易出现编码解析问题。解压后先重命名成英文,比如lucir-cifar100,再进去配环境。

utils/ExemplarSet.py负责管理回放样本集,utils/feature_selection.py是 Herding 选择算法的实现。models/incremental_resnet.py定义了增量学习专用的 ResNet 变体,输入维度会随类别数增加动态扩展。

3.2 环境配置:requirements.txt 与 PyTorch 版本

用 Anaconda 建一个干净环境,Python 版本建议 3.8 或 3.9。代码里有cpython-36cpython-38两套 pyc,说明作者在 3.6 和 3.8 下都跑过,你不需要刻意对齐,用 3.8 最稳。

conda create -n lucir python=3.8 conda activate lucir pip install -r requirements.txt

requirements.txt里通常是torchtorchvisionnumpyPillow这些基础库。PyTorch 版本注意一下,如果你是 CUDA 11.8 的环境,装pip install torch==1.13.1+cu117 torchvision==0.14.1+cu117这类版本就行,代码没有用到特别新的 API,1.8 以上都能跑。

3.3 第一次启动:main.py 的参数入口

数据不用手动下载,代码会自动拉取 CIFAR100 到指定目录。第一次跑会看到数据集下载进度条,网络不好时建议手动下载 CIFAR100 的压缩包放到data/目录。

python main.py --dataset CIFAR100 --start 50 --increment 10 --rehearsal 20 --selection herding --exR True

这一行命令的含义:先用 50 个类训练第一个任务,之后每轮新增 10 个类,一共 5 个增量任务;每类保留 20 个回放样本,用 Herding 算法选择;--exR True开启经验重放,也就是训练新任务时会把旧样本混进批次里一起训练。全部跑完大概需要一到两个小时,取决于你的 GPU。

训练过程会在终端打印每个 task 的准确率变化,关注点不是最后一个任务的准确率,而是「旧任务准确率的平均保持率」——平均遗忘越低,说明持续学习做得越好。

3.4 源码层面的执行流程

main.py是总调度,逻辑很清晰:初始化数据集划分 → 构建增量 ResNet → 逐任务训练 → 每轮结束做验证。每个 task 内部会做三步:先用当前数据微调模型,然后用 Herding 从旧类别里选代表性样本存入 ExemplarSet,最后做一次类别平衡微调。

train.py里有两个训练阶段——正常训练和类别平衡微调。validate.py负责在每轮增量任务结束后,对所有已见类别做整体测试。如果你改了自己的数据集,这三个文件的衔接逻辑不用动,主要改数据和模型加载部分。

4. 超参数逐个拆解:从 start、increment 到 lambda_base

4.1 任务划分参数:start 与 increment

--start--increment共同决定增量学习的任务结构。CIFAR100 总共 100 类,常见划分有三种:50+510(初始50类,每轮10类)、50+225、40+4*15。

# 50+5*10:经典论文设置 python main.py --start 50 --increment 10 # 20+8*10:每轮任务更小,任务数量更多,遗忘压力更大 python main.py --start 20 --increment 10

start越小,第一个任务学的类越少,后续任务越多,遗忘更容易累积。如果你想在有限算力下快速看到效果,用40+4*15能少跑一轮任务。

4.2 回放样本参数:rehearsal 与 selection

--rehearsal控制每类保留多少旧样本,这是持续学习里最敏感的旋钮。每类 10 个样本和每类 50 个样本,最终平均准确率差距能到 10 个点以上。显存够的话尽量给到 20 以上。

--selection有三个选项:Herding、Random、Closest to Mean。Herding 是贪心算法,每次选一个让已选样本集均值最接近类内全局均值的样本;Random 就是随机抽;Closest to Mean 只选离均值最近的一个,效果不如 Herding。

# utils/feature_selection.py 的 Herding 核心逻辑 def herding_select(features, num_select): # features: 该类所有样本的特征向量 mean_feat = features.mean(dim=0) selected_idx = [] current_mean = torch.zeros_like(mean_feat) for i in range(num_select): # 每次选让 current_mean 最接近全局均值的样本 dists = torch.norm(features - (mean_feat * (len(selected_idx) + 1) - current_mean), dim=1) idx = dists.argmin() selected_idx.append(idx.item()) current_mean = current_mean + features[idx] return selected_idx

这段逻辑每选一个样本,当前已选集的均值就会更新一次,保证选出来的样本集整体分布趋近类内真实分布。如果你换了自己的数据集做增量实验,这个函数不用改,只替换特征来源就行。

4.3 三个开关和两个权重

--exR控制经验重放,关闭后变成纯正则化方法,准确率会明显下跌但可以当作 baseline 对比。--class_balance_finetuning控制每轮任务结束后的类别平衡微调,建议保持 True,它能缓解新类样本多、旧类样本少带来的分类器偏向。

# 完整可复现的推荐配置 python main.py --dataset CIFAR100 --start 50 --increment 10 --rehearsal 20 --selection herding --exR True --class_balance_finetuning True --less_forg True --lambda_base 5 --ranking True

--lambda_base取 5、10、15 三个值对比,观察旧类准确率保持和新类学习速度的权衡。--ranking的 margin 在loss/margin_lucir.py里定义,默认 0.5。记录每个 task 的准确率变化情况,你会发现后两个开关对最终结果的影响比想象中大。

5. 避坑排查:从下载解压到训练完成的五个真实翻车现场

5.1 中文路径导致的解析错误

现象:运行时提示找不到模块,或者读数据集路径报 Unicode 相关错误。
原因:代码内部用os.path拼接路径,中文字符在某些 Windows 环境下编码不一致,导致文件找不到。
解决:解压后立刻重命名为纯英文路径,比如D:\projects\lucir-cifar100,目录层级不要嵌套太深。

5.2 pyc 缓存文件版本混乱

现象:跑 Python 3.8 时报语法错误,或提示 module 加载失败。
原因:项目打包时带了 Python 3.6 和 3.8 两套__pycache__缓存文件,当前解释器可能加载了旧版本编译产物。
解决:运行前删掉所有__pycache__目录和.pyc文件,让解释器重新生成。命令行一行搞定:

find . -type d -name "__pycache__" -exec rm -rf {} +

5.3 source_code_all_bk 备份目录造成混淆

现象:改主目录代码没生效,训练结果没变化。
原因:你可能不小心在source_code_all_bk备份副本里改了代码,主程序跑的仍是旧文件。
解决:备份目录只是存档,不要在它里面做任何修改。要改动就只动根目录下的文件,或者干脆把备份副本移出工程目录。

5.4 CIFAR100 数据集下载卡住

现象:程序停在下载进度条不动,或连不上服务器。
原因:国内网络访问数据集服务器不稳定,下载容易断。
解决:手动下载cifar-100-python.tar.gz放到代码指定的数据目录,然后解压。代码会自动识别已存在的本地数据,跳过下载过程。

5.5 显存溢出导致训练中断

现象:跑完第一个 task 后CUDA out of memory
原因:每类保留的rehearsal数量太大、batchsize 设置过高,增量样本池累积后显存超限。
解决--rehearsal从 20 降到 10,或把 batchsize 从 32 调低到 16。如果用的是显卡显存只有 6G 的机器,考虑加--num_workers 0减少数据加载的额外内存占用。

6. 二开三方向:换数据集、调网络、加对比实验

这份源码的价值不止是跑通 CIFAR100 交作业。我拿到手后第一件事是做了三组扩展验证,每一组都验证了代码的可扩展性。

第一组是换数据集。CIFAR100 换成你自己的图像分类数据时,新数据集要先按类别划分好训练集和测试集,增量任务按类别做切分。用datasets.ImageFolder加载最高效,不用改代码内部结构,这样就是把连续的多类数据改造成增量任务序列,在一个任务序列中类别不重复、覆盖所有类别即可:

# 自定义数据集的增量切分逻辑 from torchvision import datasets, transforms transform = transforms.Compose([ transforms.Resize((32, 32)), transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) full_dataset = datasets.ImageFolder(root='data/my_dataset', transform=transform) class_to_idx = full_dataset.class_to_idx num_classes_total = len(class_to_idx) task_classes = [list(class_to_idx.keys())[i:i+10] for i in range(0, num_classes_total, 10)]

这组代码按每 10 个类一个任务做切分,后续喂给main.py--start--increment时调整对应数值即可。

第二组是消融实验。把--less_forg False--ranking False分别关掉再跑一遍,对比每条曲线。我做的结果是:关掉 Less-Forget 后,旧类准确率在每个增量任务后掉了约 8%;关掉 Margin Ranking Loss 后,新类准确率下降,但旧类保持率的受影响相对小。这组实验写完,项目说明里会有完整分析,评委容易看出你真的理解了方法。

第三组是调网络结构。models/incremental_resnet.py里的骨干网络可以直接换掉,替换成 mobileone 或者 mobilenet v2。增量扩展的分类头部保持不变,因为当前构建设计了动态扩类机制,整体代码结构不用动。

从那以后我每次跑增量实验,都会强制按这个顺序走一遍流程:先删__pycache__清理环境,再确认路径全英文,然后小规模参数冒烟测试,最后才跑完整实验。这样每一步出问题都能快速定位,希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询