如何在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
相关产品推荐
相关产品推荐

