ARTICLE DETAIL

资讯详情

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

Triplet Loss实战指南:从三元组构造到训练避坑全流程

Triplet Loss实战指南:从三元组构造到训练避坑全流程 简介Triplet Loss三元组损失是度量学习中的重要损失函数广泛应用于人脸识别、图像检索等相似性任务。这份实战资源以MNIST手写数字数据集为场景完整给出基于Triplet Loss的模型训练与推理代码涵盖模型定义、数据加载、训练器、推理脚本等核心模块适合希望理解锚点、正负样本采样及margin设置原理的深度学习初学者或算法工程师参考。压缩包共32个文件包含20个Python脚本、8张示意图如算法流程、损失变化曲线、模型结构图、配置文件与README总大小约568KB目录结构清晰便于按模块阅读和复现。目前已有1120人学习下载。通过该资源读者可直接运行代码观察训练过程结合图示与注释理解三元组损失的计算方式、采样策略和调参思路为后续在人脸识别、图像检索等真实业务中应用打下基础。1. Triplet Loss 是什么一个看起来简单、用起来全是坑的距离度量损失在没有足够类别标签、或者类别数多到 softmax 根本扛不住的时候Triplet Loss 几乎是度量学习里的默认选项。它想做的事很直白给定一个锚点样本让模型把和它同类别的正样本拉近把不同类别的负样本推远。人脸验证、行人重识别、以图搜图、商品相似度排序这些场景里你都能看到它的影子。但真正在工程里跑过的人都会承认Triplet Loss 的收敛曲线像玄学同样的代码换个数据集效果可能天差地别。问题通常不在损失函数本身而在数据采样和 margin 设置。损失函数只是一句拉近 A、推远 B的约束但 A 和 B 怎么选才是决定模型能不能学到判别性特征的关键。这篇就走一遍完整的落地路径从构造三元组、写损失函数、跑训练循环到评估和避坑代码直接可复现。适合刚接触度量学习、想用自己的数据跑通 Triplet Loss 的工程师也适合被 loss 不下降折磨到想放弃的熟手对照排查。有一点先说在前面千万别只盯着 loss 数值Triplet Loss 的训练曲线本来就比别人难看关键在验证集上的检索效果。2. 数据是第一道坎三元组构造策略与采样代码2.1 为什么说采样比损失函数本身更重要Triplet Loss 需要的数据格式是(anchor, positive, negative)三元组。anchor 是锚点positive 和 anchor 属于同一类别negative 属于不同类别。理论上训练目标就是让 anchor 和 positive 的距离远小于 anchor 和 negative 的距离。但问题是这个远小于学到什么程度完全取决于你给模型喂了什么样的三元组。如果随便采样大部分三元组都是简单样本anchor 和 positive 本来就近negative 本来就远。模型轻轻松松就把 loss 降到很低但 embedding 空间其实没有学到什么硬性的判别能力。反过来如果只挑最难的负样本距离最近的负样本又容易让训练不稳定甚至是把模型推向一种退化状态——所有样本被挤到同一个点上loss 照样低但检索效果全无。常见做法是使用 Batch Hard 策略每个训练 batch 里包含 P 个类别、每个类别 K 张图然后在 batch 内部为每个 anchor 挑选最难的正样本和最难的正负样本。这样每次迭代都在用相对困难的样本来更新模型效率比全局随机采样高得多。实现上用 PyTorch 写一个小的采样器或者直接用torch.utils.data.Dataset配合 batch 组织逻辑两者都能跑通。2.2 用 PyTorch 构造 Batch Hard 三元组一份可以直接用的采样类先准备数据。这里用 MNIST 这类自带类别标签的数据集来演示但思路完全适用于自己的业务数据——只要每个样本有类别标签即可。下面这个采样类的输入是特征矩阵和标签向量输出是每个样本对应的 positive mask 和 negative maskimport torch def batch_hard_triplet_loss(embeddings, labels, margin0.5): embeddings: [batch_size, embed_dim] 模型输出的特征向量 labels: [batch_size] 每个样本的类别标签 margin: 正负样本距离的边界值 # 计算 batch 内所有样本两两之间的欧氏距离 # ||a - b||^2 ||a||^2 ||b||^2 - 2 * a * b dot_product torch.matmul(embeddings, embeddings.T) sq_norm torch.diag(dot_product) # 每个向量的 L2 范数平方 # 广播计算距离矩阵加上 eps 防止对角线为 0 导致除零 distance_matrix sq_norm.unsqueeze(0) sq_norm.unsqueeze(1) - 2 * dot_product distance_matrix torch.clamp(distance_matrix, min0.0) # 构造标签相等矩阵相同类别的位置为 True label_equal labels.unsqueeze(1) labels.unsqueeze(0) # 对每个 anchor找到 hardest positive同类中距离最远的 # 以及 hardest negative异类中距离最近的 # 先把对角线和异类位置置为 -inf同类位置找最大距离 distance_matrix_with_inf distance_matrix.clone() distance_matrix_with_inf[~label_equal] -float(inf) hardest_positive_dist torch.max(distance_matrix_with_inf, dim1).values # 同类位置置为 inf异类找最小距离 distance_matrix_with_inf distance_matrix.clone() distance_matrix_with_inf[label_equal] float(inf) hardest_negative_dist torch.min(distance_matrix_with_inf, dim1).values # Triplet loss: max(0, d_p - d_n margin) triplet_loss torch.clamp(hardest_positive_dist - hardest_negative_dist margin, min0.0) return triplet_loss.mean()这段代码的核心逻辑是矩阵化计算距离避免写 Python 双层循环。先把embeddings的相似度矩阵算出来再通过标签矩阵找到同类和异类的位置。构造 loss 时只取平均而不是对所有三元组求和这样数值范围稳定学习率好调。参数上需要注意几个点。margin0.5是最常见的初始值如果发现正负样本距离差距本来就很大可以适当加大到 1.0如果训练不稳定就该减小到 0.2 附近。距离用的是欧氏距离的平方因为平方后梯度形式更简单训练更平稳。如果你在做人脸验证这类场景也可以改成余弦距离只需要把 embeddings 先做 L2 归一化。距离矩阵里加了 clamp 是为了防止浮点误差产生负距离这是个小细节但没有它偶尔会算出 NaN。2.3 采样策略的取舍Batch Hard、Batch All 和 Semi-hard上面实现的是 Batch Hard它只选最难的正样本和最难负样本。Batch All 则是对 batch 内所有有效三元组求平均负样本数量多loss 更平滑但简单样本占比高收敛速度慢。Semi-hard 是只选比正样本远但差距不超过 margin的负样本介于两者之间训练最稳但实现复杂。我一般建议先用 Batch Hard 跑通因为实现简单、收敛快。如果 loss 曲线震荡太厉害再加 Semi-hard 或降低 margin。另外有个容易被忽略的点batch 的组成方式比采样策略本身更重要。一个 batch 里至少要有 8 个不同类别、每个类别 8 个样本以上否则hard样本的挑选空间太小退化成随机采样。这就是常听到的 P×K 采样P 是类别数K 是每类样本数。P16、K4 是性价比比较高的组合。3. 模型与损失函数实现embedding 网络怎么搭、Triplet Loss 怎么写3.1 损失函数的本质是约束 embedding 空间Triplet Loss 本身不关心你用 ResNet 还是 ViT它只对最后的 embedding 向量做约束。这也是它比分类损失灵活的地方分类损失要求模型输出一个类别概率分布本质上是在学决策边界Triplet Loss 是在学距离度量embedding 空间里的距离直接对应样本的相似度。所以模型结构上只需在骨干网络后面接一个 embedding 层把特征压缩到 128 维或者 256 维。128 维是检索场景的常用起点维度太低容易丢失细节太高则后续存储和检索成本上升。embedding 层要不要做归一化取决于你用什么距离。用欧氏距离可以不做用余弦距离必须先做 L2 归一化。还有一个工程细节embedding 层不要加 ReLU 激活否则输出非负会限制特征表达范围很多场景下效果会明显变差。3.2 完整可运行的 Triplet Loss 模型定义与训练骨架下面给出一个可以在 MNIST 上直接跑通的完整例子。骨干网络故意用得很简单方便你看清核心逻辑import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.data import DataLoader, Subset from torchvision import datasets, transforms import numpy as np class EmbeddingNet(nn.Module): 输出 128 维 embedding 的简单卷积网络 def __init__(self, embed_dim128): super().__init__() self.convnet nn.Sequential( nn.Conv2d(1, 32, kernel_size5, padding2), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size5, padding2), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Flatten(), nn.Linear(64 * 7 * 7, 256), nn.ReLU(inplaceTrue), ) # 最后的 embedding 层不加激活函数 self.embedding nn.Linear(256, embed_dim) def forward(self, x): return self.embedding(self.convnet(x))模型定义里有一个最常见的坑很多人会在 embedding 层后面顺手加一个tanh或ReLU以为这样可以规范化输出。实际上这会限制 embedding 的表达空间尤其是 ReLU 把负值全部截断导致大量样本的特征挤在正半轴距离度量失真。所以 embedding 层保持线性输出就好。训练循环里注意三点。第一每个 batch 的数据要手动组织成(anchor, positive, negative)的形式而不是直接把 batch 喂给模型。下面是组织逻辑def train_one_epoch(model, dataloader, optimizer, margin0.5, devicecuda): model.train() total_loss 0.0 for batch_idx, (data, labels) in enumerate(dataloader): # data: [batch_size, 1, 28, 28] # 在 batch 内为每个 anchor 随机挑选一个同类作为 positive # 随机挑选一个异类作为 negative # 更工程化的做法是像 2.2 那样直接算 batch hard loss # 这里演示的是静态三元组采样适合小数据集快速验证 # 先将数据推入 device data, labels data.to(device), labels.to(device) # 随机挑选 positive对每个样本在同类中随机选一个 positives torch.zeros_like(data) negatives torch.zeros_like(data) for i in range(data.size(0)): same_class_idx (labels labels[i]).nonzero(as_tupleTrue)[0] diff_class_idx (labels ! labels[i]).nonzero(as_tupleTrue)[0] # 排除自身选一个同类正样本 same_class_idx same_class_idx[same_class_idx ! i] if len(same_class_idx) 0: # 当前 batch 里没有同类样本跳过 continue pos_idx np.random.choice(same_class_idx.cpu().numpy()) neg_idx np.random.choice(diff_class_idx.cpu().numpy()) positives[i] data[pos_idx] negatives[i] data[neg_idx] # 前向传播得到三个 embedding anchor_emb model(data) positive_emb model(positives) negative_emb model(negatives) # 计算 triplet loss pos_dist F.pairwise_distance(anchor_emb, positive_emb, p2) neg_dist F.pairwise_distance(anchor_emb, negative_emb, p2) loss torch.mean(torch.clamp(pos_dist - neg_dist margin, min0.0)) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() return total_loss / max(len(dataloader), 1)第二静态三元组采样上面写的这种适合快速验证模型和损失函数有没有 bug但不适合完整训练。正式训练需要用 Batch Hard 策略也就是把train_one_epoch里的循环替换成调用 2.2 节写的batch_hard_triplet_loss(model(data), labels, margin)一次性算 loss不用单独构造 positive 和 negative 数据。第三反向传播时三个分支的梯度都会回传到同一个骨干网络。PyTorch 会自动累加梯度不需要手动处理。这也是注意点如果三个分支各自过了不同的 dropout 层会破坏距离语义所以模型里尽量别在 embedding 层前加 dropout。如果想要正则化不如加 weight decay。3.3 margin 和距离度量的选择逻辑margin 是 Triplet Loss 里最值得花时间调的参数。它的含义是正样本对距离必须比负样本对距离小至少 margin 这么多才不算产生 loss。margin 太小模型学到微小的距离差就觉得满足了embedding 空间区分度差margin 太大模型被逼着把正样本对压缩到非常近、负样本对推到非常远容易把 embedding 空间撑到体积无限大训练长时间无法收敛。经验取值是 0.2 到 1.0 之间。距离用欧氏距离L2时margin 取 0.5 起步比较稳妥用余弦相似度时margin 要在 0.1 到 0.5 之间调因为余弦相似度本身有界。如果你同时跑多个模型做对比实验记得所有模型用同一种距离和同一个 margin否则对比没有意义。4. 训练循环与调参从 loss 曲线到反向传播的完整链路4.1 联合训练为什么交叉熵损失能帮 Triplet Loss 一把Triplet Loss 有个众所周知的毛病训练初期 embedding 空间还没成形hard sample 的选择基本等于随机挑模型很难稳定起步。常见解决方案是联合训练——在分类头比如 softmax 交叉熵损失和 Triplet Loss 之间做一个加权和。分类头在一开始主导训练方向让 embedding 至少具备基本的类别可分性等分类 loss 开始下降后Triplet Loss 再慢慢接管精修距离度量。也和你做实验用的损失函数曲线图有关。联合训练时两个 loss 不要放在同一个坐标轴里看它们的量级完全不同。交叉熵量级一般在个位数Triplet Loss 在 0.1 量级混在一起根本看不出趋势。训练时分别记录两个 loss画两条曲线。如果 Triplet Loss 波动很大说明采样策略和 margin 得改如果分类 loss 一直在降但 Triplet Loss 不动说明 margin 设置太大几乎没有三元组能产生 loss。4.2 一个完整的训练脚本集成 Batch Hard、联合损失、学习率调度import torch from torch.utils.data import DataLoader, RandomSampler from torchvision import datasets, transforms # 使用 MNIST 做演示换成自己的数据只需替换 dataset 部分 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) full_dataset datasets.MNIST(./data, trainTrue, downloadTrue, transformtransform) # 取前 5000 个样本加快实验迭代 train_dataset Subset(full_dataset, range(5000)) train_loader DataLoader(train_dataset, batch_size128, shuffleTrue, num_workers2) model EmbeddingNet(embed_dim128).cuda() # 分类头只用于联合训练不参与最终的 embedding 提取 classifier nn.Linear(128, 10).cuda() # 两个优化器分开设置学习率embedding 网络用 1e-3分类头用 1e-3 optimizer torch.optim.Adam([ {params: model.parameters()}, {params: classifier.parameters()} ], lr1e-3, weight_decay5e-4) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size10, gamma0.5) triplet_margin 0.5 lambda_triplet 0.5 # Triplet loss 的权重联合训练时常用 0.1 ~ 1.0 for epoch in range(30): model.train() classifier.train() total_triplet_loss 0.0 total_ce_loss 0.0 for data, labels in train_loader: data, labels data.cuda(), labels.cuda() embeddings model(data) # 分支 1Triplet LossBatch Hard t_loss batch_hard_triplet_loss(embeddings, labels, margintriplet_margin) # 分支 2交叉熵分类损失 logits classifier(embeddings) ce_loss F.cross_entropy(logits, labels) # 联合训练反向传播 loss lambda_triplet * t_loss ce_loss optimizer.zero_grad() loss.backward() optimizer.step() total_triplet_loss t_loss.item() * data.size(0) total_ce_loss ce_loss.item() * data.size(0) scheduler.step() avg_triplet total_triplet_loss / len(train_dataset) avg_ce total_ce_loss / len(train_dataset) print(fEpoch {epoch1:02d} | Triplet Loss: {avg_triplet:.4f} | CE Loss: {avg_ce:.4f})联合训练时最核心的旋钮是lambda_triplet。它控制在总 loss 里 Triplet Loss 的占比。0.5 是折中值如果你的数据类别多但每个类样本少建议降到 0.2 以下让交叉熵先尽力把大类分开Triplet Loss 只做微调。反过来如果你已经有一个在相关任务上预训练好的模型lambda_triplet可以提到 1.0——因为 embedding 已经有先验结构不需要交叉熵带路。优化器的选择上Adam 是要快于 SGD 的。Triplet Loss 的梯度噪声本来就大SGD 的收敛速度会慢到让人怀疑模型有 bug。weight decay 给到 5e-4 是经验值。学习率调度用 StepLR 就够每 10 个 epoch 减半。注意别用 ReduceLROnPlateau它需要监控验证集 loss而 Triplet Loss 的验证指标不适合直接用 loss 来判断——验证集上 loss 高不一定效果差反而说明模型在努力把难样本拉开。4.3 训练时看什么指标loss 之外要盯的距离分布单独画 Triplet Loss 曲线会让你误判。比如 loss 在 1.0 附近震荡你以为是没收敛实际上正负样本距离分布已经重叠度很低了。所以训练时除了打印 loss还要周期性统计两个距离分布的平均值。实现很简单在验证集上算 anchor-positive 距离均值和 anchor-negative 距离均值两个值之差就是 margin 的实际余量。如果之差超过 marginloss 就会趋近于 0但 embedding 空间可能仍然不好——而如果之差是负数说明正样本平均距离反而比负样本远这个模型是彻底的废了。def evaluate_distance_distribution(model, dataloader, devicecuda): 在验证集上统计正样本对和负样本对的平均距离 model.eval() pos_dists [] neg_dists [] with torch.no_grad(): for data, labels in dataloader: data, labels data.to(device), labels.to(device) embeddings model(data) dist_matrix torch.cdist(embeddings, embeddings, p2) labels_eq labels.unsqueeze(1) labels.unsqueeze(0) # 对角线排除 mask torch.eye(data.size(0), dtypetorch.bool).to(device) pos_mask labels_eq ~mask neg_mask ~labels_eq ~mask if pos_mask.sum() 0: pos_dists.append(dist_matrix[pos_mask].mean().item()) if neg_mask.sum() 0: neg_dists.append(dist_matrix[neg_mask].mean().item()) return np.mean(pos_dists), np.mean(neg_dists)这个函数的输出非常有信息量。假设pos_dist0.8, neg_dist0.9说明虽然 loss 很低但正负样本在 embedding 空间里只拉开了 0.1 的距离检索结果几乎不可用。反过来pos_dist0.3, neg_dist1.2差距有 0.9这才是一个健康的 embedding 空间。我一般在每个 epoch 末尾打印这两个值比单独看 loss 可靠得多。5. Triplet Loss 避坑记录五条血泪踩坑实况5.1 现象loss 下降很快但检索效果很差这是最迷惑人的情况。loss 从 1.2 降到 0.05 只用了 2 个 epoch你以为模型学得很好一测 Recall1 只有 30%。原因是 batch 里简单三元组太多。随机采样时大部分负样本离 anchor 很远Triplet Loss 很快就学会了把 easy negative 推开实际没有学到精细的判别特征。解决方法是换成 Batch Hard 采样或者至少保证 batch 里类别足够多P≥16迫使模型处理难负样本。5.2 现象训练中期 loss 突然暴涨甚至出现 NaN这个情况在 Batch Hard 里经常出现。某个 batch 里有一个难到极致的负样本距离接近 0margin - neg_dist变成很大的正数梯度爆炸。原因有两个一是数据里存在错误标签把同类的样本标成了负类二是 embedding 没有做数值稳定处理。解决办法先在数据层面做标签清洗再用距离裁剪。把 loss 计算改为torch.clamp(pos_dist - neg_dist margin, min0.0, max10.0)限制每个三元组对总 loss 的贡献上限。max 值取 10 是在限制梯度噪声同时不牺牲正常 hard sample 的学习。5.3 现象模型输出所有样本的同一条 embedding这种崩溃叫 embedding collapse。所有样本的特征向量都变成同一个常数向量距离全是 0Triplet Loss 为 0看似完美收敛实则没有任何信息。常见诱因是 margin 设得过大加上 batch 里的 hard negative 被选得太极端梯度把 embedding 往一个点压缩。另一个诱因是用了 Batch All 策略平均了太多简单三元组的梯度模型无法从困难样本上获得有效信号。应对方法是先把 margin 降到 0.2然后在训练里随机丢弃 20% 的三元组给梯度注入随机性。如果还是不恢复就在 embedding 层后面接一个 BN 层强制每一维度的分布不过度集中。5.4 现象验证集 loss 正常但训练集 loss 几乎为 0这种情况通常不是过拟合而是你用了静态三元组采样训练时提前固定了 anchor/positive/negative模型把三条路径都记住了。之前训练骨架里演示的静态采样就是为了快速验证用的正式训练千万别用。解决办法是确保每个 epoch 重新采样三元组。动态采样有两个方案一是在 DataLoader 的__getitem__里每次随机选 positive 和 negative二是按 Batch Hard 那样在 batch 内动态选。前者简单但每次得到的三元组质量波动大后者按困难度选训练效率高。5.5 现象不同类别样本数量差距大少数类完全学不出来Triplet Loss 对类别均衡的要求比交叉熵损失更高。如果某个类别只有 2 个样本在一个 batch 里很难和别的样本组成有效三元组这个类别的 embedding 基本靠随机初始化撑着。基础解决法是类别平衡采样每个类别每个 epoch 内至少出现 K 次K 等于 batch size 除以类别数。进阶做法是类内增强对少数类样本做随机裁剪、翻转、色彩抖动先把训练样本数做上去。我自己做过的项目里用后一种方法把少数类的 Recall1 从 42% 拉到了 61%效果显著。6. 评估与进阶用 RecallK 和距离可视化验证 embedding 质量6.1 RecallK 评估代码训练完模型最终极的一步是验证 embedding 质量。这里不能用分类准确率因为 Triplet Loss 的目标不是分类而是距离度量所以要用检索指标给定一个 query在 gallery 里找最相似的 K 个样本看其中有没有和 query 同类的。def evaluate_recall_at_k(model, gallery_loader, query_loader, k10): gallery: 候选池每个样本有 embedding 和标签 query: 查询样本在 gallery 中检索同类别样本 model.eval() gallery_embs [] gallery_labels [] with torch.no_grad(): for data, labels in gallery_loader: data data.cuda() emb model(data) gallery_embs.append(emb.cpu()) gallery_labels.extend(labels.tolist()) gallery_embs torch.cat(gallery_embs, dim0) # [N, 128] hits 0 total 0 with torch.no_grad(): for data, labels in query_loader: data data.cuda() query_emb model(data).cpu() # [B, 128] # 计算 query 到所有 gallery 样本的距离 dist torch.cdist(query_emb, gallery_embs, p2) # [B, N] # 排除 query 本身就是 gallery 的情况跨数据集评估时可以忽略 _, topk_idx dist.topk(k, largestFalse, dim1) for i in range(query_emb.size(0)): query_label labels[i].item() retrieved_labels [gallery_labels[idx] for idx in topk_idx[i].tolist()] if query_label in retrieved_labels: hits 1 total 1 recall_at_k hits / max(total, 1) return recall_at_k使用上注意两点。第一gallery 和 query 最好来自不同的人或不同的采集批次否则你在做的是记住训练数据而不是泛化检索。第二topk 检索结果里如果包含 query 自身同一个样本会虚高 Recall。跨数据集评估天然避免了这个问题但如果只用单数据集切分需要把包含 query 本身的结果从 gallery 里移掉或者让 gallery 全部来自验证集。6.2 用 t-SNE 可视化最直接的 embedding 质量证据数值指标之外我强烈建议做一次 t-SNE 可视化把验证集的 embedding 向量降维到 2D 画散点图。这一步能让你一眼看出三个问题不同类别的点有没有聚成团、有没有类别重叠、有没有类别被压成一条线。数据准备和画图代码可以直接交给 sklearn 的TSNEfrom sklearn.manifold import TSNE import matplotlib.pyplot as plt def visualize_embedding(model, dataloader, num_samples2000, save_pathembedding_tsne.png): model.eval() embs [] labels [] with torch.no_grad(): for data, label in dataloader: data data.cuda() embs.append(model(data).cpu()) labels.extend(label.tolist()) if len(embs) * data.size(0) num_samples: break embs torch.cat(embs, dim0)[:num_samples] labels labels[:num_samples] tsne TSNE(n_components2, random_state42, perplexity30) embs_2d tsne.fit_transform(embs.numpy()) plt.figure(figsize(8, 6)) scatter plt.scatter(embs_2d[:, 0], embs_2d[:, 1], clabels, cmaptab10, s10, alpha0.7) plt.colorbar(scatter) plt.title(Embedding Visualization with t-SNE) plt.savefig(save_path, dpi150)稳定复现的关键参数是perplexity30。数据量超过 3000 个样本时要调大 perplexity比如 50否则局部结构会失真。观察可视化结果时如果发现同一类别的点分散成多个小簇说明模型没有学到类别内的紧致性可以考虑增加 margin 或者调大特征维度。6.3 最后一个技巧距离阈值校准检索场景里你往往不只是要 top-K而是要一个相似/不相似的判定阈值。比如人脸闸机比对得分超过某阈值才放行。这个阈值千万不要凭感觉设。正确做法是在验证集上画出正样本对距离分布和负样本对距离分布的直方图取两个分布交叉点作为初始阈值再根据业务对误识率和拒识率的要求做微调。这个阈值校准我在多个项目里吃过亏。曾经有个人脸巡检项目验证集 EER 曲线算下来最优阈值是 0.82我图省事用了 0.75上线后误识别率直接翻了倍。后来养成的习惯是每次模型更新必须重新跑一次距离分布统计阈值跟着变它是个活参数不是定死值。希望帮到你——按这套流程把你的 Triplet Loss 项目从头推到尾loss 曲线再难看也能拿出让人信服的检索效果来。本文还有配套的精品资源点击获取
返回列表