ARTICLE DETAIL

资讯详情

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

基于 Triplet Loss 的相似度学习实战:从 FashionMNIST 到距离度量式图像区分

基于 Triplet Loss 的相似度学习实战:从 FashionMNIST 到距离度量式图像区分 基于 Triplet Loss 的相似度学习实战从 FashionMNIST 到距离度量式图像区分【免费下载链接】visionDatasets, Transforms and Models specific to Computer Vision项目地址: https://gitcode.com/gh_mirrors/vi/vision相似度学习Similarity Learning的目标是训练模型输出一个可度量的嵌入向量embedding使得同类别样本在嵌入空间内彼此靠近、不同类别样本彼此远离。本文以 references/similarity 参考实现为核心完整讲解基于三元组损失Triplet Loss的嵌入学习全流程损失函数与两种三元组挖掘策略的源码原理、保证批次类别结构的PKSampler采样器、ResNet50骨干的嵌入网络设计以及一份开箱即用、可在 10 个 epoch 内于 FashionMNIST 测试集上复现 97% 准确率的训练脚本。读完本文你将能够理解并改造这套代码把三元组损失应用到类别数量未知或需要学习距离度量如人脸相似度的业务场景。背景为什么用三元组损失做相似度学习传统图像分类要求预先定义固定的类别集合输出层维度等于类别数模型学到的是这张图属于哪一类的判别边界。而相似度学习解决的是另一类问题模型输出一个固定维度的嵌入向量类别边界由嵌入之间的几何距离隐式定义。这一学习方式由 FaceNet: A Unified Embedding for Face Recognition and ClusteringOlivia Huang 等2015推广开来在人脸识别与聚类任务中被证明极其有效。其核心思想是不直接预测身份而是学习一个映射使得同一身份的人脸嵌入向量之间的距离小于不同身份人脸之间的距离。本参考实现正是这一思想在 torchvision 生态中的落地它直接适用于以下两类场景类别数量未知你无法预先枚举所有类别但仍希望模型学会区分它们例如新增用户、新增商品无需重训分类头需要距离度量你的业务本质是衡量样本之间的相似程度例如做人脸相似度比对、以距离作为检索或去重的排序依据。核心概念Anchor、Positive 与 Negative 三元组TripletMarginLoss的输入不是单个样本或样本对而是一个三元组triplet包含三个样本角色含义要求Anchor锚点一个取自数据集的样本作为参照基准Positive正样本与 Anchor 属于同一标签/分组的一个样本通常要求positive ! anchor即二者是不同的样本Negative负样本与 Anchor 属于不同标签/分组的一个样本标签与 Anchor 不同训练目标非常直观同类别样本的嵌入应该彼此近不同类别样本的嵌入应该彼此远。损失函数正是把这一目标转化为可微分的数学表达驱动网络通过梯度下降实现它。TripletMarginLoss损失函数的数学定义与源码实现数学定义README 中给出了TripletMarginLoss的核心公式loss max(dist(anchor, positive) - dist(anchor, negative) margin, 0)其中dist是距离函数默认使用欧几里得距离即 L2 距离。最小化该损失会同时产生两个效果最小化dist(anchor, positive)把同类别样本在嵌入空间中拉近最大化dist(anchor, negative)把不同类别样本推开。margin是间隔超参数它规定正负样本距离差至少要达到多少才认为区分得足够好。当dist(anchor, positive)已经比dist(anchor, negative)小出至少margin时损失为 0网络不再对该三元组产生梯度。源码实现仓库中的 loss.py 提供了一个继承自torch.nn.Module的TripletMarginLoss封装其构造与调用方式为class TripletMarginLoss(nn.Module): def __init__(self, margin1.0, p2.0, miningbatch_all): super().__init__() self.margin margin self.p p self.mining mining if mining batch_all: self.loss_fn batch_all_triplet_loss if mining batch_hard: self.loss_fn batch_hard_triplet_loss def forward(self, embeddings, labels): return self.loss_fn(labels, embeddings, self.margin, self.p)三个关键参数margin间隔默认 1.0训练脚本中通过--margin设为 0.2p距离函数使用的范数阶数p2.0即欧几里得距离mining三元组挖掘策略可选batch_all或batch_hard下文详述。值得注意的是forward的入参是整个批次的嵌入向量与标签而非预先组好的三元组——三元组的构造与筛选在损失函数内部完成这是该实现与朴素逐三元组计算的重要区别。三元组挖掘策略batch_all 与 batch_hard从一批样本中穷举所有合法三元组的计算量随批次大小呈三次方增长且大量三元组损失为 0、对训练毫无贡献。因此 README 强调在把三元组送入损失函数之前需要先进行挖掘mining。本实现支持两种策略二者通常都能加速训练收敛。batch_all全量三元组 困难样本筛选对应 loss.py 中的batch_all_triplet_loss流程如下用torch.cdist计算批次内两两样本的成对距离矩阵通过广播构造所有(anchor, positive, negative)组合的损失张量anchor_positive_dist - anchor_negative_dist margin用_get_triplet_mask生成的掩码剔除不合法组合——要求三个样本索引互不相同且 Anchor 与 Positive 标签相同、Anchor 与 Negative 标签不同见 loss.py将损失小于 0 的易样本easy triplet置零剔除对剩余损失大于 0 的困难三元组求均值并额外返回正三元组比例fraction_positive_triplets作为训练日志指标。该策略的优点是信息利用充分缺点是计算量较大、受批次中类别与样本数量影响明显。batch_hard为每个锚点挑选最难三元组对应 loss.py 中的batch_hard_triplet_loss思路更为激进最难正样本hardest positive对每个 Anchor在其所有同标签样本中选取距离最远的那个anchor_positive_dist.max(1)最难负样本hardest negative对每个 Anchor在其所有异标签样本中选取距离最近的那个其中非法负样本通过累加行内最大距离被掩盖见 loss.py每个 Anchor 只产生一个三元组最终对全部 Anchor 的三元组损失求均值第二个返回值固定为-1表示不适用。相比batch_allbatch_hard每个 Anchor 只贡献一个梯度信号计算效率更高且梯度始终由最难区分的样本驱动训练信号更强。两种策略的取舍从源码结构可以推断实际使用时需要根据批次大小与数据情况权衡batch_all更稳适合探索阶段全面利用批次信息batch_hard更狠收敛通常更快但若批次内负样本整体距离过近可能带来较大的方差。二者的梯度质量都高度依赖批次的类别结构——这正是下一个组件PKSampler存在的意义。PKSampler保证批次内类别结构的采样器为什么普通 DataLoader 不够TripletMarginLoss需要在同一个批次内找到同标签样本对Anchor-Positive与异标签样本对Anchor-Negative。如果批次由随机采样得到同一标签往往只出现 12 个样本合法三元组会极度稀疏挖掘策略形同虚设。为此sampler.py 实现了PKSampler每个批次大小为p * k恰好包含p个类别的样本每类恰好k个样本。分组与合法性校验构造PKSampler(groups, p, k)时首先调用create_groups把数据集样本索引按标签分组并剔除样本数不足k的组见 sampler.py确保每类都能凑齐k个样本。随后校验剩余组数是否不少于p否则抛出ValueError(There are not enough classes to sample from)见 sampler.py。采样流程__iter__的核心逻辑见 sampler.py对每个组内样本索引做一次随机打乱记录每组剩余可用样本数当剩余可采样的组数大于p时循环随机选出p个组用torch.multinomial均匀抽样从每个组中按顺序取出k个样本组内已打乱无需再随机并产出某个组剩余样本数不足k后将其移出候选组集合避免产出残缺批次。这个打乱后顺序取用 耗尽即淘汰的设计保证每个批次都严格满足p类 ×k样本的结构同时不重复使用同一组内已产出的样本。单元测试验证tests 侧对应的 test.py 用torchvision.datasets.FakeData对采样器做了两个关键断言当p16大于数据集的类别数10时构造PKSampler会抛出断言错误阻止非法配置遍历完整数据加载器后每个批次的类别集合大小恰好为p且每个类别恰好出现k个样本。这段测试从代码层面印证了PKSampler对批次结构的硬性保证是复用该采样器时的行为契约。EmbeddingNet模型定义与 L2 归一化model.py 定义了嵌入网络class EmbeddingNet(nn.Module): def __init__(self, backboneNone): super().__init__() if backbone is None: backbone models.resnet50(num_classes128) self.backbone backbone def forward(self, x): x self.backbone(x) x nn.functional.normalize(x, dim1) return x要点默认骨干为torchvision.models.resnet50并把最后的全连接层输出维度改为128即每个样本映射到 128 维嵌入向量forward末尾对嵌入做L2 归一化nn.functional.normalize(x, dim1)把嵌入约束到单位超球面上。归一化后欧几里得距离与余弦相似度直接对应且距离有明确上界便于后续用固定阈值做判定骨干网络可替换构造时传入任意 backbone如models.resnet18即可这为控制模型体积与推理速度留出了改造空间。训练脚本完整参数说明与数据管线训练脚本 是整套参考实现的可运行入口。默认配置下它在FashionMNIST数据集上训练ResNet50学习可用于区分图像的嵌入README 明确说明按默认参数运行可在10 个 epoch 内于 FashionMNIST 测试集上达到约 97% 的准确率。运行方式python train.py -h # 列出全部可选参数 python train.py # 使用默认参数运行训练命令行参数一览以下参数定义见 train.py各参数含义、默认值与取值范围如下参数默认值说明--dataset-dir/tmp/fmnist/FashionMNIST 数据集存放目录首次运行自动下载-p, --labels-per-batch8每个批次包含的唯一类别数-k, --samples-per-label8每个类别在批次中的样本数批次大小 p * k默认 64--eval-batch-size512评估阶段的批次大小--epochs10总训练轮数-j, --workers4数据加载的工作进程数--lr0.0001Adam 优化器的初始学习率--margin0.2三元组损失的间隔超参数--print-freq20每多少个 batch 打印一次训练日志--save-dir.模型权重保存目录--resume待加载的检查点路径用于断点续训--test-onlyFalse仅加载权重并评估不训练需配合--resume--use-deterministic-algorithmsFalse强制使用确定性算法保证结果可复现其中p、k是影响训练质量最直接的两个参数p越大单个批次内可构造的困难三元组越多样k越大每个类别的正样本对越充足。二者的乘积即为批次大小受显存约束。数据预处理管线输入图像经过如下变换见 train.pytransforms.Compose([ transforms.Lambda(lambda image: image.convert(RGB)), # 灰度图转 RGB适配 ResNet 三通道输入 transforms.Resize((224, 224)), # 缩放到 ResNet 标准输入尺寸 transforms.PILToTensor(), # PIL 图像转张量 transforms.ConvertImageDtype(torch.float), # 归一化到 [0, 1] 浮点 ])FashionMNIST 本身是 28×28 灰度图而 ResNet50 需要 224×224 的三通道输入因此前三步完成了灰度→RGB、小图→大图的适配ConvertImageDtype把像素值从uint8转为[0, 1]的浮点张量。训练循环与日志训练主流程 的关键设计训练集使用DataLoader(train_dataset, batch_sizep*k, samplerPKSampler(targets, p, k), ...)其中targets是形如第 i 个样本的标签的列表PKSampler依赖它完成分组采样每个 step取样本 → 前向得到嵌入 →criterion(embeddings, targets)返回(loss, frac_pos_triplets)→ 反向传播 → Adam 更新每print_freq个 batch 打印一次平均损失与困难三元组百分比% avg hard triplets该指标由batch_all策略返回的正三元组比例换算而来可用于观察批次内困难样本的稀疏程度每个 epoch 结束后评估并保存权重文件名形如epoch_n__ckpt.pth--resume加载检查点时使用torch.load(args.resume, weights_onlyTrue)符合新版 PyTorch 的安全加载要求。评估方法距离矩阵与最佳阈值搜索相似度学习的评估与分类不同没有现成的输出层对比标签可做。本实现采用成对距离判定见 train.py在torch.inference_mode()下用model.eval()推理整个测试集收集全部嵌入与标签用torch.cdist计算嵌入两两之间的完整距离矩阵构造标签是否相同的目标矩阵取距离矩阵的上三角不含对角线作为样本对在0.01到1.50、步长0.01的候选阈值上逐一扫描find_best_threshold见 train.py判定距离 ≤ 阈值为同一类别统计与真实标签一致的样本对比例输出使准确率最高的阈值与对应的最高准确率例如accuracy: 97.000%, threshold: 0.85。这种最佳阈值 准确率的评估方式与推理阶段完全一致——实际部署时只需把两幅图像的嵌入距离与训练好的阈值比较即可判断是否属于同一类别。完整复现步骤在具备 PyTorch 与 torchvision 环境的机器上按以下步骤即可端到端复现# 1. 进入参考实现目录 cd references/similarity # 2. 查看所有可选参数 python train.py -h # 3. 使用默认参数训练约 10 个 epoch测试集准确率约 97% python train.py # 4. 指定数据集目录、批次结构与学习率进行定制训练 python train.py --dataset-dir /data/fmnist -p 16 -k 8 --lr 0.0001 --epochs 15 --save-dir ./ckpt # 5. 仅评估已有检查点 python train.py --resume ./ckpt/epoch_10__ckpt.pth --test-only运行环境要求需要可用的 CUDA 设备脚本默认torch.device(cuda:0 if torch.cuda.is_available() else cpu)无 GPU 时自动回退 CPU但 224×224 的 ResNet50 在 CPU 上训练耗时显著增加如需保证结果可复现可加--use-deterministic-algorithms。扩展到自定义数据集与真实业务README 明确指出训练脚本的默认配置可以按需修改从源码结构看迁移到自己的数据只需改动三处替换数据集将FashionMNIST换成任意分类数据集如torchvision.datasets.ImageFolder。PKSampler只要求提供targets第 i 个样本标签的列表ImageFolder 自带同名属性可直接使用见 train.py 的注释说明调整预处理把Resize((224, 224))换成目标骨干所需的输入尺寸RGB 转换与归一化照常保留替换骨干与嵌入维度在 model.py 中向EmbeddingNet传入自定义 backbone并把num_classes改成业务所需的嵌入维度。在类别数量未知、需要距离度量的业务中如人脸相似度比对、以图搜图、重复样本去重训练完成的嵌入网络即可投入使用计算查询样本与库中样本的嵌入欧几里得距离按训练阶段确定的阈值判定是否同类。参考资料FaceNet 论文详细描述了三元组损失的动机与细节Olivier Moindrot 关于 triplet loss 的系列文章对采样与挖掘策略有系统讲解本仓库的 loss.py 即是其 PyTorch 移植模块 docstring 中明确标注了出处。这两个外部资料均可作为理解本文所述概念特别是batch_all与batch_hard背后的数学动机的延伸阅读而本文涉及的实现细节均以上述仓库源码为准。【免费下载链接】visionDatasets, Transforms and Models specific to Computer Vision项目地址: https://gitcode.com/gh_mirrors/vi/vision创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表