// contrastive_loss.mbt -- Metric learning losses (v0.130.0).
//
// Three canonical embedding-space losses:
//
// 1. Contrastive loss (Bromley et al. 1993; Hadsell et al. 2006):
// L = (1 - y) * d^2 + y * max(0, margin^2 - d^2)
// where y = 1 for a dissimilar pair and y = 0 for a similar pair,
// and d^2 is the squared distance between the two embeddings.
//
// 2. Triplet loss (FaceNet, Schroff et al. 2015):
// L = max(0, d(a, p)^2 - d(a, n)^2 + margin)
// for an anchor a, a positive p (same class) and a negative n.
//
// 3. N-pair loss (Soh et al. 2016): a softmax cross-entropy over the
// similarity matrix, which is numerically better behaved than
// triplet loss when mining hard negatives.
//
// Scope of v0.130.0:
// - contrastive_pair_loss
// - triplet_loss
// - n_pair_loss
// - hard_negative_triplet_loss (max over candidate negatives)
//
// Reference: Bromley 1993; Hadsell 2006; FaceNet 2015; N-pair 2016.
///|
/// Contrastive loss on one pair of embeddings. `similar` is 1.0 for a
/// same-class pair (should be close) and 0.0 for a different-class pair
/// (should be far).
pub fn contrastive_pair_loss(
a : Array[Float],
b : Array[Float],
similar : Float,
margin : Float,
) -> Float {
let d2 = pairwise_squared_distance(a, b)
if similar > 0.5F {
// Dissimilar pair: push apart until d^2 > margin^2.
let m2 = margin * margin
if d2 < m2 {
m2 - d2
} else {
0.0F
}
} else {
// Similar pair: pull together.
d2
}
}
///|
/// Triplet loss: max(0, d(a,p)^2 - d(a,n)^2 + margin).
pub fn triplet_loss(
anchor : Array[Float],
positive : Array[Float],
negative : Array[Float],
margin : Float,
) -> Float {
let d_pos = pairwise_squared_distance(anchor, positive)
let d_neg = pairwise_squared_distance(anchor, negative)
let l = d_pos - d_neg + margin
if l > 0.0F {
l
} else {
0.0F
}
}
///|
/// Hard-negative triplet loss: take the maximum loss over a set of
/// candidate negatives. This is online hard-negative mining (FaceNet
/// section 3.1) and is what makes triplet training work in practice.
pub fn hard_negative_triplet_loss(
anchor : Array[Float],
positive : Array[Float],
negatives : Array[Array[Float]],
margin : Float,
) -> Float {
let m = negatives.length()
if m == 0 {
return 0.0F
}
let mut worst = 0.0F
for i in 0.. worst {
worst = l
}
}
worst
}
///|
/// N-pair loss (Soh et al. 2016). Given an anchor and `num_classes`
/// positive/negative pairs (one positive + one negative per class),
/// computes the softmax cross-entropy over the similarity logits:
///
/// logits_j = -2 * ||a - p_j||^2 - ||a - n_j||^2
/// L = -log( softmax(logits)_correct )
///
/// Negative similarity logits are offset by a constant (the paper uses
/// 0.1) for numerical stability, which is a known wart of the N-pair
/// formulation.
pub fn n_pair_loss(
anchor : Array[Float],
positives : Array[Array[Float]],
negatives : Array[Array[Float]],
) -> Float {
let k = positives.length()
if k == 0 {
return 0.0F
}
let logits : Array[Float] = Array::make(k, 0.0F)
let mut m = 0.0F
for j in 0.. m {
m = s
}
}
// Stable softmax; the correct class is index 0.
let mut sum_exp = 0.0F
for j in 0..