尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

如何训练NetVLAD: trainWeakly弱监督三元组损失与困难负样本挖掘全流程图解

如何训练NetVLAD: trainWeakly弱监督三元组损失与困难负样本挖掘全流程图解 如何训练NetVLAD: trainWeakly弱监督三元组损失与困难负样本挖掘全流程图解【免费下载链接】netvladNetVLAD: CNN architecture for weakly supervised place recognition项目地址: https://gitcode.com/gh_mirrors/ne/netvladNetVLAD 是面向弱监督地点识别place recognition的经典开源 MATLAB 项目它让 CNN 仅凭图片GPS 坐标这样的粗略标签就能学会把同一地点的图片聚到一起。本文围绕核心训练脚本trainWeakly.m用图解方式拆解三元组损失与**困难负样本挖掘Hard Negative Mining**的完整流程并附上关键参数速查表帮助新手快速上手训练自己的地点识别网络。1️⃣ 什么是弱监督为什么不需要逐张标注传统图像识别需要给每张图片打精确标签而 NetVLAD 的弱监督范式只需要图片的地理位置坐标数据集基类 datasets/dbBase.m 中保存了每张数据库图和查询图的 UTM 坐标utmDb、utmQ以及一个距离阈值posDistThr两张图片只要距离小于阈值就自动互为正样本无需人工标注训练时再通过nontrivialPosQ方法筛选非平凡正样本距离 1~nonTrivPosDistSqThr米避免用同一全景图的近邻造成作弊式学习。内置数据集如dbTokyoTimeMachine东京时间机器训练/验证、dbTokyo247东京 24/7测试、dbPitts匹兹堡都在 datasets/ 目录下想用自己的数据只需继承dbBase编写一个轻量规格文件即可。2️⃣ 训练前准备三步环境清单 ✅依赖MATLAB relja_matlab MatConvNet≥ v1.0-beta18强烈建议安装 Yael_matlab 加速最近邻搜索仓库内 yael_dummy/ 提供了yael_nn、yael_kmeans的纯 MATLAB 简化实现训练慢很多但零门槛。预训练网络从 MatConvNet 官网下载imagenet-caffe-refAlexNet或imagenet-vgg-verydeep-16VGG-16路径配置把localPaths.m.setup复制为localPaths.m填好依赖、数据集、预训练模型的路径然后在 MATLAB 中运行setup;。 新手可以先跑demo.m文件末尾的dbTiny 微型训练示例——几分钟即可跑完用来验证环境是否配置正确。3️⃣ trainWeakly 训练六步全流程 训练入口极其简洁以 Tokyo Time Machine VGG-16 为例dbTrain dbTokyoTimeMachine(train); dbVal dbTokyoTimeMachine(val); sessionID trainWeakly(dbTrain, dbVal, ... netID, vd16, layerName, conv5_3, backPropToLayer, conv5_1, ... method, vlad_preL2_intra, ... learningRate, 0.0001, doDraw, true);trainWeakly.m内部按以下六步执行Step 1加载并裁剪主干网络loadNet(vd16, conv5_3)加载 VGG-16 并裁到最后一个卷积层conv5_3因为 NetVLAD 需要一个卷积特征图作为输入。Step 2自动加装聚合层addLayersaddLayers.m解析method字符串并逐段装配网络头方法片段作用vlad/vladv2自动用 kmeansk64见getClusters.m在训练描述子上聚类生成可学习 VLAD 层preL2聚合前对特征逐维做 L2 归一化intraVLAD 输出做 intra-归一化每个聚类块内 L2 归一化max/avg备选的最大/平均池化基线末尾固定postL2整向量 L2 归一化layerWholeL2Normalize.mStep 3特征缓存serialAllFeats用初始网络一次性算出所有训练库图/查询图及验证集的图像表示存为二进制.bin文件serialAllFeats.m。后续挖掘困难负样本时直接查这个缓存而不是逐张过网络。Step 4基线测试test0训练前先对缓存特征算一次 RecallN记录未训练off-the-shelf的基线方便日后对比收益。Step 5主训练循环核心每个 epoch 随机打乱查询顺序按batchSize4 个三元组迭代每个查询执行挖掘→前向→反向三部曲详见下一节。Step 6周期性保存与验证每saveFrequency个 batch 和每个 epoch 结束保存 checkpoint文件名含sessionID与 epoch 号每 epoch 用computeAllFeats重算特征并跑testNet把训练/验证 RecallN 与 rankloss 写入obj学习率每lrDownFreq个 epoch 除以lrDownFactor同时把特征重算频率compFeatsFrequency乘上同倍数因为网络学得慢后缓存更耐放。4️⃣ 核心难点三元组损失 困难负样本挖掘 这是trainWeakly.m约 L264~L413的灵魂所在。对每个查询 q流程如下查询 q │ ├─① 找最近正样本 p在 nontrivialPosQ 候选里用 yael_nn 求 dPos │ ├─② 组装负样本候选池 上轮记忆(nNegCache10) ∪ 随机抽样(nNegChoice1000) │ ├─③ 距离筛选violate d(q,n)² dPos margin │ └─ 只保留最多 nNegCap10 个最难的负样本 n │ ├─④ 前向把 [q, p, n₁…nₖ] 一起过网络拿到真实特征缓存只是用来选样本 │ ├─⑤ 损失L Σ max(d(q,p) margin − d(q,n), 0) │ └─⑥ 反向手写梯度公式 → vl_simplenn 反传 → SGD 动量更新权重几个关键设计点margin默认 0.1正样本距离与负样本距离之间的安全间隔只有违反 margin的负样本才产生梯度困难负样本记忆nNegCache每个查询把当前最难的 10 个负样本存进auxData.negCache跨 epoch 保留。下一轮它们往往是顽固错误样本重点回炉二次验证先用缓存特征粗筛前向拿到真实特征后重新判定 violatingL358~L372防止用过期缓存误导梯度excludeVeryHard默认 false若设为 true则丢弃比正样本还近的超难样本只保留半难样本semi-hard。原作者实验未开启但如果你打算从头训练而非微调预训练网络可以参考 FaceNet 的思路试试。5️⃣ 关键参数速查表 以下默认值来自 trainWeakly.m 头部调优优先级从高到低margin / 负样本数 → 学习率组 → 缓存频率。参数默认值含义netID/layerNamecaffe/conv5主干网络与截断层VGG 用vd16/conv5_3methodvlad_preL2_intra聚合方式NetVLAD 官方推荐配置backPropToLayer1全网络梯度反传到的最深层如conv5_1可加速且防过拟合learningRate0.0001初始学习率lrDownFreq/lrDownFactor5/2每 5 个 epoch 学习率减半batchSize4每个 SGD batch 的三元组数一个三元组含多张图实际显存占用更大margin0.1三元组 marginnNegChoice1000每轮随机抽样的负样本数nNegCap10每个三元组实际保留的困难负样本上限nNegCache10跨 epoch 记忆的困难负样本数nEpoch30训练轮数compFeatsFrequency1000每 1000 个查询重算一次特征缓存saveFrequency2000保存 checkpoint 的间隔查询数epochTestFrequency1建议保持 1否则pickBestNet无法工作6️⃣ 避坑指南曲线怎么看、问题怎么查 ⚠️打开doDrawtrue或事后用plotResults可以看到三类曲线来自obj结构动态损失左上训练 batch 上的平滑三元组损失。⚠️ 由于困难负样本挖掘会不断出难题损失不降甚至小幅上升是正常现象真正该盯的是 recall 曲线rankloss左下在固定采样三元组上评估的损失可跨 epoch 对比RecallN右侧最重要的曲线t.5表示训练集 recall5v.5表示验证集。常见故障对照症状原因与对策动态损失出现周期性尖峰特征缓存更新太慢网络对缓存过拟合 →调小compFeatsFrequency反之若无问题且想加速可调大多进程训练冲突/文件损坏相同数据集网络方法的进程会写同一个缓存文件 → 给每个进程设不同checkpoint0suffix中途手动停止后重跑报错输出.bin文件不完整 →删除损坏文件再重跑已存在的文件不会重算想断点续训传startEpoch 对应sessionID即可7️⃣ 训练收尾挑出最佳网络 PCA 白化 训练结束后或中途崩溃后用验证集表现挑出最优 epoch 的网络[~, bestNet] pickBestNet(sessionID); % 默认按 val recall5 挑选 finalNet addPCA(bestNet, dbTrain, doWhite, true, pcaDim, 4096);pickBestNet.m内部调用getBestEpoch做带 tie-breaking 的最优轮次选择并打印训练后 vs 未训练的 recall 对比addPCA配合白化能显著压缩维度、降低存储同时提升地点识别精度。最后照搬 README 的测试四步曲构造dbTokyo247→serialAllFeats算库图/查询图特征 →testFromFn计算 RecallN → 画图即可得到完整的检索性能曲线。一句话总结NetVLAD 的训练 预训练 CNN 自动聚类的 NetVLAD 层靠trainWeakly.m中缓存特征选样、真实特征求梯度的三元组损失循环不断优化配合 margin 与三层负样本控制choice/cap/cache实现高效的困难负样本挖掘——理解这条主线再调参实验就水到渠成了 【免费下载链接】netvladNetVLAD: CNN architecture for weakly supervised place recognition项目地址: https://gitcode.com/gh_mirrors/ne/netvlad创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表