简介:面向深度学习中度量学习与相似度匹配任务,一份完整工程代码与配套数据以 MNIST 手写数字集为场景,覆盖从损失函数原理到训练推理落地的全流程。包内共 32 个文件、约 568KB,其中 20 个 Python 脚本按 data_loaders、models、trainers、utils、configs 等模块组织,便于快速定位数据采样、模型构建、训练器与推理入口;另有 8 张训练曲线、算法流程等示意图像,以及 JSON 配置文件、requirements 依赖清单和 README 说明文档,可完整支撑项目复现。内容不仅包含 Triplet Loss 的公式与 margin 设置、三种典型三元组采样策略,还给出了 MNIST 上的实际训练与推理效果对比,帮助理解图像检索和相似性度量的核心思路。已有 1120 人学习浏览,代码注释与结构化目录对入门玩家友好,适合想快速跑通损失函数实战的开发者。
1. Triplet Loss 损失函数:为什么分类头训得再好,检索任务照样翻车
做行人重识别或者人脸验证的朋友大概率都遇过这种尴尬:把分类头训到99%的准确率,一上测试集,遇到没见过的ID照样懵。分类头学的是“这个人是几号人”,一旦类别集合变了,输出层就得重训。Triplet Loss 损失函数的思路完全相反:它不看身份,只看距离,让同一人的特征靠近,不同人的特征远离。这个方向在实际项目里的坑比想象中多——margin怎么设、采样怎么采、loss降到多少算收敛,每一步都是玄学。这篇就按我落地时的顺序,把三元组怎么构造、代码怎么写、距离怎么评估讲明白,新手能顺着跑通,熟手直接翻到第5章看避坑和第6章的进阶技巧。
2. 三元组与损失公式:搞懂 anchor、positive、negative 之间的关系,margin 不是拍脑袋定的
2.1 为什么说 Triplet Loss 解决的是“开放集”问题,而不是“分类”问题
交叉熵配合softmax分类头解决的是封闭集问题:训练时见过所有类别,测试时也只在这些类别里判断。但人脸验证、商品检索、ReID这类任务,测试时出现的Identity往往是训练时根本没见过的。拿分类头去提取特征,特征空间里根本没有“这个新人和谁靠近”的约束,结果就是检索列表里前排全是无关样本。
Triplet Loss把学习目标从“分类正确”改成“距离正确”。每次从训练数据里抽三个样本:一个anchor(锚点),一个positive(和anchor同类的正样本),一个negative(和anchor不同类的负样本)。损失函数要求anchor与positive的距离尽可能近,anchor与negative的距离尽可能远,并且远到一定程度之外。这样学出来的embedding天然适合做最近邻检索——新样本进来,不用分类,直接算特征距离就能找到相似的旧样本。
这个思路最直观的应用就是把一张查询图和一个候选库里的图都过一遍网络,取embedding算余弦相似度或欧氏距离,排序输出。整个过程没有“类别”概念,所以类别在训练后被新增、合并、删除都不影响使用。这一点是分类头替代不了的。
2.2 公式拆解:d(a,p) 与 d(a,n) 的差值就是网络要优化的全部
Triplet Loss的标准形式是:
loss = max(d(a,p) - d(a,n) + margin, 0)
其中 d(a,p) 是anchor与positive的特征距离,d(a,n) 是anchor与negative的特征距离,margin 是一个需要手工设定的超参数。直观理解:网络希望 d(a,p) 比 d(a,n) 小至少 margin 这么多;如果已经满足了,loss为0,梯度不更新;如果不满足,loss为正值,梯度推动网络把正样本拉近、负样本推远。
当d(a,p)已经等于0时,loss = max(-d(a,n) + margin, 0),这时负样本距离只要超过margin,loss就为0。注意这里的距离空间由特征归一化方式决定:如果对embedding做了L2归一化,欧氏距离的取值范围是[0, 2],margin就不能设得太大,超过2的margin在数学上永远无法满足,loss永远不会为0;如果不做归一化,embedding的模长可以随意增长,margin就需要按实际距离尺度估算。很多第一次上手的人在这里翻车——margin设了1.0,配合L2归一化,理论上d范围只有[0,2],看似可行,但实际上大部分负样本对的距离集中在1.2~1.8之间,需要拉开0.6~1.0的差距非常困难,训练过程会一直震荡。
梯度方面也值得看一眼。对a、p、n三个输入的梯度方向不同:a往远离n、靠近p的方向同时移动,p只往靠近a的方向移动,n只往远离a的方向移动。如果只更新a而不更新p和n,收敛会慢得多,所以实际训练时三元组的三个样本都要参与反向传播,batch内样本利用率也会因此变高。
2.3 采样策略决定了训练效率:random 采样为什么经常不收敛
损失函数本身很简单,真正让Triplet Loss难训的是采样策略。随机从数据集里抽三元组,往往抽到的是easy triplets——anchor和negative的距离已经非常远,loss为0,梯度为0,网络学不到东西。训练集可能已经过了几十个epoch,loss却一直在一个低位徘徊,embedding质量也没提升,问题很可能就出在采样上。
业界常见的三种策略:
- Offline triplet mining:每个epoch前先把所有样本过一遍网络,算出所有距离,离线挑出满足条件的hard triplets,再喂给网络训练。缺点是每个epoch都要额外推理一次,计算开销大,而且离线算的距离在模型更新后立刻过期。
- Online triplet mining:在训练过程中,从当前batch内部构造三元组。因为batch里的特征都是最新模型算出来的,不会有过期问题,这是最常用的方案。
- BatchHard策略:从每个batch里为每个anchor挑选最难的正样本(距离最大的同类)和最难的负样本(距离最小的异类)。这个策略迫使模型处理最难的情况,收敛快,但容易受噪声标签影响——如果某个正样本标签标错了,它会被当成“最难的”反复强化,模型就被带偏了。
我一般用online + BatchHard,因为实现简单、训练效率高。噪声问题靠数据清洗和限制hard程度来缓解,比如只选top-K最难的负样本,而不是选最难的一个。
3. 用 PyTorch 实现 Triplet Loss:数据加载、BatchHard 策略、损失函数模块,代码可复制可跑
3.1 数据准备:构造一个能产出三元组的 Dataset 类
先从数据说起。标题里写了“完整代码+数据”,最省事的数据集是MNIST或Fashion-MNIST,torchvision自带下载逻辑,不需要额外手动准备文件,改一行代码就能替换成自己的数据文件。下面这个Dataset类接受一个普通的分类数据集(每个样本是图片和类别ID),在__getitem__里动态生成三元组。
import torch from torch.utils.data import Dataset import random class TripletMNIST(Dataset): """ 从普通分类数据集中构造三元组。 samples: list of (image_tensor, label) """ def __init__(self, samples): self.samples = samples self.labels = [s[1] for s in samples] # 建一个 label -> 样本索引列表 的映射,方便采正样本和负样本 self.label_to_indices = {} for idx, label in enumerate(self.labels): self.label_to_indices.setdefault(label, []).append(idx) def __len__(self): return len(self.samples) def __getitem__(self, idx): anchor_img, anchor_label = self.samples[idx] # 正样本:从同一类里随机挑一个,不能是anchor自己,否则距离恒为0 pos_indices = [i for i in self.label_to_indices[anchor_label] if i != idx] if len(pos_indices) == 0: # 如果这个类别只有一个样本,退而求其次用anchor本身(训练时影响有限) pos_idx = idx else: pos_idx = random.choice(pos_indices) # 负样本:从所有其他类别里随机挑一个 neg_label = random.choice( [l for l in self.label_to_indices.keys() if l != anchor_label] ) neg_idx = random.choice(self.label_to_indices[neg_label]) pos_img = self.samples[pos_idx][0] neg_img = self.samples[neg_idx][0] return anchor_img, pos_img, neg_img, anchor_label逻辑说明:这个类每次返回一个三元组,anchor由外部索引决定,positive从同一标签中随机选取,negative从不同标签中随机选取。注意当某个类别只有一个样本时,pos_idx退化成idx本身,d(a,p)恒为0,这个三元组对训练几乎没有贡献。因此实际使用时要保证每个类别至少有两个以上的样本,或者重采样时尽量保证类别均衡。
参数说明:这里用随机采样构造三元组,运行速度快但容易产生easy triplets,所以后续把三元组交给BatchHard模块再筛一遍。如果只想跑通最小demo,这个随机版本就够用;如果要做正式实验,我建议把3.2的BatchHard直接接上。
3.2 BatchHard 策略:如何在 batch 内自动挖掘最难的样本
常见的做法是在一个batch里同时放入P个类别、每个类别K张图,对每张图作为anchor时,在该batch内找到距离最大的正样本和距离最小的负样本。下面用全距离矩阵实现,输入是batch的embedding矩阵,输出是loss。
import torch import torch.nn as nn class BatchHardTripletLoss(nn.Module): """ 输入: embeddings (B, D), labels (B,) 输出: 标量loss """ def __init__(self, margin=0.3): super().__init__() self.margin = margin def forward(self, embeddings, labels): # 1. 计算成对欧氏距离矩阵 # ||a-b||^2 = ||a||^2 + ||b||^2 - 2*a·b dot = torch.mm(embeddings, embeddings.t()) sq_norm = torch.diag(dot) dist_sq = sq_norm.unsqueeze(0) + sq_norm.unsqueeze(1) - 2 * dot dist_sq = torch.clamp(dist_sq, min=0.0) dist = torch.sqrt(dist_sq + 1e-9) # 加极小值防止梯度在0处断裂 # 2. 构造mask: 同类和异类 labels_eq = labels.unsqueeze(0) == labels.unsqueeze(1) # (B, B) # 3. 对每个anchor,找出难正样本(同类中距离最大)和难负样本(异类中距离最小) batch_size = embeddings.size(0) hardest_positive = torch.zeros(batch_size, device=embeddings.device) hardest_negative = torch.zeros(batch_size, device=embeddings.device) for i in range(batch_size): pos_dist = dist[i][labels_eq[i]] neg_dist = dist[i][~labels_eq[i]] if pos_dist.numel() > 0: hardest_positive[i] = pos_dist.max() if neg_dist.numel() > 0: hardest_negative[i] = neg_dist.min() # 4. 计算triplet loss loss = torch.clamp(hardest_positive - hardest_negative + self.margin, min=0.0) return loss.mean()逻辑说明:成对距离矩阵的计算用了一个常见的展开技巧,避免显式循环嵌套。label_eq矩阵标出所有同类位置。对每个anchor单独取同类中的最大距离和异类中的最小距离,得到hardest positive和hardest negative,再套max(d(p)-d(n)+margin, 0)后取均值。如果batch里某个anchor没有正样本或负样本,对应位置保持0,不参与loss,实际使用时避免在小batch里出现这种情况。
参数说明:margin是核心超参,通常从0.1、0.2、0.3这几个值开始试;embedding如果做L2归一化,margin建议不超过1.0。batch_size要尽量大,常见做法是P×K,比如8个类别每类8张图,batch_size=64,这样才能保证每个anchor都能找到有意义的难负样本。batch太小的时候,异类距离可能全局都很大,hardest negative也构不成多大挑战,训练效果会明显变差。
3.3 完整训练脚本:从原始数据到能用的embedding模型
把数据加载、模型、损失函数拼起来,跑一个可以直接观察损失函数曲线图的最小训练脚本。
import torch import torch.nn as nn import torch.optim as optim from torchvision import datasets, transforms from torch.utils.data import DataLoader # 1. 数据准备 transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) mnist_train = datasets.MNIST(root='./data', train=True, download=True, transform=transform) # 转成list方便三元组Dataset使用 samples = [(img, label) for img, label in mnist_train] # 2. 简单CNN特征提取网络 class EmbeddingNet(nn.Module): def __init__(self, embedding_dim=64): super().__init__() self.features = nn.Sequential( nn.Conv2d(1, 32, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2), nn.Flatten(), nn.Linear(64 * 7 * 7, embedding_dim), ) # 对输出做L2归一化,让距离尺度稳定在[0,2]内 self.normalize = nn.functional.normalize def forward(self, x): x = self.features(x) return self.normalize(x, dim=1) # 3. 组装训练 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = EmbeddingNet(embedding_dim=64).to(device) triplet_loss_fn = BatchHardTripletLoss(margin=0.3) optimizer = optim.Adam(model.parameters(), lr=1e-3) # 这里为了演示用随机采样的Dataset,正式训练建议按P×K方式组织batch dataset = TripletMNIST(samples[:20000]) loader = DataLoader(dataset, batch_size=64, shuffle=True, num_workers=2) loss_history = [] for epoch in range(10): epoch_loss = 0.0 for anchor, positive, negative, _ in loader: anchor, positive, negative = anchor.to(device), positive.to(device), negative.to(device) optimizer.zero_grad() # 分别过网络取embedding emb_a = model(anchor) emb_p = model(positive) emb_n = model(negative) # 拼成一个大batch,方便BatchHard计算 embeddings = torch.cat([emb_a, emb_p, emb_n], dim=0) labels = torch.cat([ torch.arange(anchor.size(0), device=device), torch.arange(anchor.size(0), device=device), torch.arange(anchor.size(0) + 1000, anchor.size(0) + 1000 + anchor.size(0), device=device) ]) # 上面这个labels有问题,见下方说明 loss = triplet_loss_fn(embeddings, labels) loss.backward() optimizer.step() epoch_loss += loss.item() avg_loss = epoch_loss / len(loader) loss_history.append(avg_loss) print(f"epoch {epoch+1}, loss: {avg_loss:.4f}")这里必须说明一个我故意挖的坑:上面labels的构造方式不对。TripletMNIST返回的三元组是独立采样的,anchor、positive、negative三个batch之间没有对齐关系,直接拼接起来用BatchHard计算,labels无法正确表达谁和谁是同类。这是新手最容易犯的错误——把Triplet Dataset的输出直接喂给BatchHard。
正确的组织方式是P×K采样:一个batch包含P个ID、每个ID有K张图,把所有图同时过网络,然后用这批图自己的labels算BatchHard。上面为了篇幅简化了逻辑,正式跑的时候建议直接用下面的Sampler思路,或者放弃BatchHard,用三元组Dataset + 普通TripletLoss(直接计算每组a/p/n的距离),那样虽然训练慢但逻辑不会出错。
参数说明:embedding_dim=64在MNIST这种简单数据集上够用,换到人脸/商品数据集常见128或256。learning_rate为1e-3配Adam是通用起步值,训练时观察损失函数曲线图,如果震荡明显可以降到3e-4。normalize操作让embedding落在单位超球面上,这样欧氏距离和余弦相似度只是单调变换关系,检索排序结果一致。
4. 训练与评估:损失函数曲线图怎么解读,embedding 质量怎么量化,三个关键参数怎么调
4.1 用损失函数曲线图判断训练状态:loss 为 0 不一定是好事
训练过程中把每个epoch的loss记录下来,画损失函数曲线图是最直接的监控手段。这里有个反直觉的现象:Triplet Loss收敛到0不代表embedding好用,只代表当前batch内所有anchor的最难正样本距离都比最难负样本距离小至少margin。如果采样太easy,模型根本不需要学到好的特征,loss也能降到接近0。
判断模型是否真正学到了可泛化的距离关系,要看三个信号:训练集loss下降曲线的斜率是否平稳;在验证集上手动构造hard三元组算loss,看是否同步下降;以及直接跑检索评估,看Recall@K是否提升。第三个信号最可靠。我习惯每两个epoch保存一次模型,在验证集上算一次Recall@1,同时把loss曲线和Recall曲线画在同一张图里对比,如果loss还在降但Recall已经停滞,往往是过拟合到训练集的hard样本模式上了,这时要加大数据增强或换更难的采样策略。
用matplotlib画曲线图就三行代码,把上面训练脚本里记录的loss_history直接传进去:
import matplotlib.pyplot as plt plt.plot(range(1, len(loss_history) + 1), loss_history, marker='o') plt.xlabel('epoch') plt.ylabel('triplet loss') plt.title('Triplet Loss Training Curve') plt.grid(True) plt.show()逻辑说明:横轴epoch、纵轴loss,直观看出收敛趋势。配合验证集Recall曲线一起看才有意义。曲线前期快速下降、后期平缓属于正常;如果中期出现回升,大概率是学习率偏大或者采样策略导致梯度不稳定。
4.2 用 Recall@K 和距离分布评估 embedding:不要只看 loss
检索类任务的标准评估指标是Recall@K:对每个查询样本,在候选集中找出与它特征距离最近的K个样本,如果其中至少有一个和它同类的样本,就算命中。下面给一个不加库依赖的PyTorch实现,输入是全部样本的embedding矩阵和标签,输出Recall@1、@5等。
def recall_at_k(embeddings, labels, k=5): """ embeddings: (N, D) 已经归一化 labels: (N,) """ embeddings = torch.nn.functional.normalize(embeddings, dim=1) # 距离矩阵 dist = torch.cdist(embeddings, embeddings) N = labels.size(0) correct = 0 for i in range(N): # 排除自己,取最近的K个 knn_idx = dist[i].argsort()[1:k+1] if labels[knn_idx].eq(labels[i]).any(): correct += 1 return correct / N # 用法示例:把验证集所有样本过模型,取embedding后计算 model.eval() with torch.no_grad(): all_embs, all_labels = [], [] for imgs, lbls in val_loader: imgs = imgs.to(device) all_embs.append(model(imgs).cpu()) all_labels.append(lbls) all_embs = torch.cat(all_embs) all_labels = torch.cat(all_labels) print("Recall@1:", recall_at_k(all_embs, all_labels, k=1)) print("Recall@5:", recall_at_k(all_embs, all_labels, k=5))逻辑说明:torch.cdist计算所有样本两两间欧氏距离。argsort取最近的K个索引时从1开始,因为索引0是自身,距离为0。Recall@1对embedding质量最敏感,因为要求最近邻必须同类;Recall@5更宽容,适合类别多、类内差异大的场景。
参数说明:评估时要注意候选集不要包含查询样本本身,否则会把“自己匹配自己”当成命中,指标虚高。工业级做法是把查询集和候选集分开,查询样本不在候选集中出现;如果只有一个全量库,就要排除自己,就像上面代码里从索引1开始取。
4.3 三个必调参数与一组可抄的起始配置
Triplet Loss训练里最影响结果的三个参数:
- margin:控制正负样本对之间的距离差要求。设得太大,模型始终学不到“足够好”,loss高居不下;设得太小,网络稍微拉开一点距离就觉得满足了,embedding区分度不够。常见做法是先设0.2~0.3跑一版,看验证集Recall再朝两个方向各试一组。
- batch_size:Triplet Loss对batch_size的敏感度远超分类任务。batch越大,batch内难负样本越难,梯度信息量越大。GPU显存允许的情况下尽量大,我常在ReID任务上用P=16、K=4,batch_size=64起步。
- embedding维度:维度太低容纳不下细粒度差异,维度太高容易过拟合且检索存储成本高。人脸任务常见128/256,商品检索64/128,MNIST这类简单数据64就够。
一张可以直接照抄的起始配置表,按这个跑通后再逐步调:
| 参数 | 建议值 | 调整方向 |
|---|---|---|
| margin | 0.3 | loss不降就调小,检索结果不分开就调大 |
| batch组织 | P=8, K=8 | 增加P提高负样本多样性 |
| embedding维度 | 64 | 简单数据64,复杂数据128~256 |
| 优化器 | Adam | 稳定;换SGD需要更仔细调lr |
| 学习率 | 1e-3 | 曲线震荡就降到3e-4 |
| 特征归一化 | L2归一化 | 配合欧氏距离,让距离范围稳定 |
5. 避坑与排错:Triplet Loss 训练翻车的五条血泪经验
5.1 现象:loss 长期不下降,徘徊在0.5~0.8
原因分析:margin设得太大,或者采样到的三元组大多数是hard negative不够hard但也没有简单到loss为0,模型一直在“拉开距离但永远拉不到margin要求”的状态。另一个常见原因是学习率过低,模型更新幅度太小。
解决方法:先把margin降到0.2跑50个epoch看曲线斜率。如果loss还在高位,打印一批d(a,p)和d(a,n)的实际数值分布,确认当前平均差距;比如d(a,p)均值0.8、d(a,n)均值0.6,说明负样本比正样本还近,模型确实没学到,此时检查batch组织是否出现大量错标数据,再考虑把学习率提高到3e-3。
5.2 现象:loss 很快就降到接近 0,但验证集 Recall@1 只有 30%
原因分析:采样太easy。随机采样的三元组大部分是easy triplets,d(a,p)已经远小于d(a,n),loss为0,模型没有收到有效梯度。loss降为0只是假象,embedding并没有把难样本分开。
解决方法:切换到BatchHard策略,强制每个batch都使用最难三元组。如果BatchHard后loss从0回升,说明模型之前确实没有见过hard样本,这是正确的训练状态。另一种做法是离线挖掘hard三元组,每轮训练前用当前模型重新采样一次。注意BatchHard的batch内部要有足够的类内样本,否则找不到有意义的hardest positive。
5.3 现象:训练过程 loss 正常下降,但同一样本在相邻两个 epoch 的特征距离突变
原因分析:没有固定数据增强管线,或者BatchNorm的running statistics在embedding网络里剧烈波动。Triplet Loss对特征分布极敏感,BatchNorm在batch较小时统计量不稳定,导致每轮输出距离尺度漂移。
解决方法:embedding网络最后一层不要接BatchNorm,或者在推理时锁定running statistics;训练时固定数据增强的随机种子,保证同一个样本在不同epoch经过的增强操作可复现性更强。还有一个细节,嵌入层输出后的L2归一化要在训练和推理时保持一致,有的实现只在推理时归一化,训练时不归一化,效果就会不稳定。
5.4 现象:batch_size 调到 128 以上直接 OOM
原因分析:BatchHard需要计算B×B的距离矩阵,显存占用是O(B²)而不是O(B)。B=128时距离矩阵就有16384个浮点数,再加上embedding的中间激活,显存消耗远高于同等batch_size的分类网络。
解决方法:限制batch_size,比如P=12、K=4,batch_size=48,距离矩阵只有2304个项。如果确实需要大batch,可以分块计算距离矩阵,每次只算一个batch块和另一个batch块的距离,累积难样本索引后再算loss。也可以用梯度累积模拟大batch,但要注意BatchHard内部的难样本挖掘只在一个子batch内进行,和真正的全局大batch效果不完全一样。
5.5 现象:验证集距离分布显示同类样本还没有完全聚拢,但训练集 loss 已经很低
原因分析:出现过拟合。Triplet Loss的难样本挖掘会把模型往“区分训练集里的难样本”方向推,如果训练集本身噪声大、或某个类别图片数量过少,模型学到的是训练集特有的模式,而不是通用的类别语义特征。
解决方法:训练时增加数据增强——随机裁剪、颜色抖动、旋转对MNIST影响不大,但对人脸/商品图效果明显;降低embedding维度,减少模型过拟合空间;或者把难样本挖掘从“最难的一个”改成“最难的几个取平均”,降低噪声样本对梯度的主导作用。另一个实用做法是在验证集上做K折交叉验证,确认Recall指标的提升不是某一折的偶然现象。
提示:Triplet Loss的调试核心是先看距离分布、再看loss曲线、最后看Recall。只盯着loss数值调参,大概率会调进死胡同。
6. 进阶技巧:把固定 margin 改成自适应 cosine margin,顺手验证一下特征分布
上面所有代码用的都是欧氏距离加L2归一化,相当于余弦距离的单调映射。但在商品检索、人脸验证这些场景里,直接优化余弦距离常常比欧氏距离更稳。原因是欧氏距离在归一化后小角度变化对距离影响不敏感,而余弦形式可以让模型更关注方向差异。一个具体改法是把固定的margin替换成随夹角变化的cosine margin:
loss = max(cos(a,p) - cos(a,n) + margin, 0) 当 cos(a,p) - cos(a,n) 与原始公式方向相反时需要取负号
写成PyTorch的自定义损失函数:
class CosineTripletLoss(nn.Module): def __init__(self, margin=0.3): super().__init__() self.margin = margin def forward(self, embeddings, labels): # 已经L2归一化的情况下,余弦相似度 = 点积 sim = torch.mm(embeddings, embeddings.t()) # B, B bs = embeddings.size(0) pos_sim = torch.zeros(bs, device=embeddings.device) neg_sim = torch.zeros(bs, device=embeddings.device) eq = labels.unsqueeze(0) == labels.unsqueeze(1) for i in range(bs): if eq[i].sum() > 1: pos_sim[i] = sim[i][eq[i] & (torch.arange(bs, device=embeddings.device) != i)].max() neg_sim[i] = sim[i][~eq[i]].min() # 我们希望pos_sim高、neg_sim低;loss为负时截断为0 loss = torch.clamp(neg_sim - pos_sim + self.margin, min=0.0) return loss.mean()逻辑说明:余弦相似度越大代表越接近,所以损失函数的比较方向相对于欧氏距离反过来了:希望neg_sim - pos_sim + margin尽量小。前面用torch.arange构造一个非对角线mask的写法有点绕,但避免了把样本自己和自己的相似度当正样本的常见错误。margin建议在0.2~0.4之间起步,因为归一化后的余弦相似度取值范围是[-1, 1],margin超过1就失去意义。
配套的验证手段:训练结束后把所有验证集样本的embedding投影到二维,或者直接画距离矩阵热力图。如果同类的距离明显小于异类距离,热力图应该呈现清楚的块状结构——对角线附近的块是同类,颜色深(距离近),其他位置浅。看不到块状结构就说明embedding还没学好,回去调采样策略而不是继续加epoch。这个技巧我每次换新数据集都会跑一次,两分钟就能直观判断模型有没有学到想要的东西,比只看loss曲线可靠得多。
希望帮到你。
本文还有配套的精品资源,点击获取