You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在TensorFlow的Triplet Loss中使用余弦相似度

TensorFlow下基于余弦相似度的三元组损失实现

三元组损失(Triplet Loss)的基础定义如下:

L(A, P, N) = max(‖f(A) - f(P)‖² - ‖f(A) - f(N)‖² + margin, 0)

公式中涉及的三类样本与参数含义:

  • A=anchor:锚点样本
  • P=positive:和锚点属于同一类别的正样本
  • N=negative:和锚点属于不同类别的负样本
  • margin:锚点与正样本、负样本之间需要满足的最小距离间隔阈值

常规三元组损失默认使用L2距离作为度量方式,实际场景中可以用(1 - cosine_similarity)替换L2距离完成距离度量。
特别注意:TensorFlow框架内置的余弦相似度损失取值范围为[-1,1],和通用余弦相似度指标的取值逻辑完全相反:取值为0时表示两个向量正交,取值越接近-1代表两个向量相似度越高,取值越接近1代表两个向量差异越大。

已验证可用的实现代码

经过实际测试,以下基于TensorFlow内置余弦相似度损失实现的三元组损失可以正常运行:

self.margin = 1
self.loss = tf.keras.losses.CosineSimilarity(axis=1)
# 计算锚点与正样本的距离、锚点与负样本的距离
ap_distance = self.loss(anchor, positive)
an_distance = self.loss(anchor, negative)
loss = tf.maximum(ap_distance - an_distance + self.margin, 0.0)

后续也可以尝试使用TensorFlow Addons内置的三元组损失接口,只需要对齐对应接口要求的数据传入格式即可。

单独调用余弦相似度计算的参考写法

如果不需要封装成损失函数,只是单独计算两个向量的余弦相似度,可以参考以下实现:

cosine_similarity = tf.keras.metrics.CosineSimilarity()
cosine_similarity.reset_state()
cosine_similarity.update_state(anch_prediction, other_prediction)
similarity = cosine_similarity.result().numpy()

参考资料

  • PyTorch CosineEmbeddingLoss 层官方文档
  • TensorFlow Keras 内置余弦相似度损失源码实现
  • TensorFlow Addons 硬间隔/软间隔三元组损失官方文档

内容的提问来源于stack exchange,提问作者Edv Beq

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.28 04:48:16