尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

LUCIR增量学习实战:基于余弦归一化与特征约束缓解灾难性遗忘

发布时间:2026/9/24 18:49:02

资讯中心
01
ARTICLE

LUCIR增量学习实战:基于余弦归一化与特征约束缓解灾难性遗忘

LUCIR增量学习实战:基于余弦归一化与特征约束缓解灾难性遗忘
简介面向图像分类持续学习任务的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要解决的核心就是灾难性遗忘——神经网络在新数据上做梯度更新时旧任务的特征表征会被逐渐覆盖。这份基于 LUCIRLarge 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, p2, dim1) masked_loss (l2_loss * old_classes_mask).sum() / old_classes_mask.sum() return masked_lossold_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, margin0.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 的 505*10 划分再分别关闭--less_forg和--ranking做消融能直观看到每个组件对最终准确率的贡献。3. 把工程跑起来目录结构、环境与第一次训练3.1 解压后的目录全景下载解压后你会看到两层结构根目录是主工程source_code_all_bk是备份副本。主目录里main.py是入口train.py和validate.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-36和cpython-38两套 pyc说明作者在 3.6 和 3.8 下都跑过你不需要刻意对齐用 3.8 最稳。conda create -n lucir python3.8 conda activate lucir pip install -r requirements.txtrequirements.txt里通常是torch、torchvision、numpy、Pillow这些基础库。PyTorch 版本注意一下如果你是 CUDA 11.8 的环境装pip install torch1.13.1cu117 torchvision0.14.1cu117这类版本就行代码没有用到特别新的 API1.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_base4.1 任务划分参数start 与 increment--start和--increment共同决定增量学习的任务结构。CIFAR100 总共 100 类常见划分有三种50510初始50类每轮10类、50225、404*15。# 505*10经典论文设置 python main.py --start 50 --increment 10 # 208*10每轮任务更小任务数量更多遗忘压力更大 python main.py --start 20 --increment 10start越小第一个任务学的类越少后续任务越多遗忘更容易累积。如果你想在有限算力下快速看到效果用404*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(dim0) 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), dim1) 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(rootdata/my_dataset, transformtransform) class_to_idx full_dataset.class_to_idx num_classes_total len(class_to_idx) task_classes [list(class_to_idx.keys())[i:i10] 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__清理环境再确认路径全英文然后小规模参数冒烟测试最后才跑完整实验。这样每一步出问题都能快速定位希望帮到你。本文还有配套的精品资源点击获取
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。