// siamese_net.mbt -- SiameseNet with shared weights (v0.131.0).
//
// A Siamese network (Bromley et al. 1993) runs the *same* embedding
// function over two or three inputs and compares the resulting
// embeddings. Weight sharing is what makes the learned metric
// symmetric and reduces the parameter count by a factor of the number
// of inputs.
//
// This file provides the twin-architecture wrapper around
// `EmbeddingNet` (v0.129.0):
//
//   anchor   \
//             >-- SiameseNet (shared EmbeddingNet) --> distances
//   positive /
//   negative /
//
// Scope of v0.131.0:
//   - SiameseNet struct: one shared EmbeddingNet + margin + type.
//   - SiameseNet::new.
//   - siamese_embed: embed one input.
//   - siamese_embed_pair: embed two inputs through the same weights.
//   - siamese_embed_triplet: embed three inputs through the same weights.
//   - siamese_distances: compute the pair/triplet distance matrix.
//   - siamese_num_params (shared weights counted ONCE).

///|
/// Loss type selected for a SiameseNet. Contrastive is pair-based;
/// Triplet is anchor/positive/negative based.
pub(all) enum SiameseLoss {
  Contrastive
  Triplet
} derive(Eq, Debug)

///|
/// SiameseNet: a shared EmbeddingNet plus the inference margin.
pub struct SiameseNet {
  net : EmbeddingNet
  margin : Float
  loss_type : SiameseLoss
}

///|
/// Build a fresh SiameseNet.
pub fn SiameseNet::new(
  net : EmbeddingNet,
  margin : Float,
  loss_type : SiameseLoss,
) -> SiameseNet {
  { net, margin, loss_type }
}

///|
/// Embed a single input through the shared weights.
pub fn siamese_embed(s : SiameseNet, x : Array[Float]) -> Array[Float] {
  embedding_net_forward(s.net, x)
}

///|
/// Embed two inputs through the same weights and return both
/// embeddings. This is the "twin" in siamese.
pub fn siamese_embed_pair(
  s : SiameseNet,
  x1 : Array[Float],
  x2 : Array[Float],
) -> (Array[Float], Array[Float]) {
  let e1 = embedding_net_forward(s.net, x1)
  let e2 = embedding_net_forward(s.net, x2)
  (e1, e2)
}

///|
/// Embed three inputs through the same weights and return all three
/// embeddings. This is the "triplet" variant.
pub fn siamese_embed_triplet(
  s : SiameseNet,
  a : Array[Float],
  p : Array[Float],
  n : Array[Float],
) -> (Array[Float], Array[Float], Array[Float]) {
  let ea = embedding_net_forward(s.net, a)
  let ep = embedding_net_forward(s.net, p)
  let en = embedding_net_forward(s.net, n)
  (ea, ep, en)
}

///|
/// Compute the pairwise distance matrix (in embedding space) for a
/// batch of `n` already-embedded vectors. Returns a flat row-major
/// [n x n] array of squared distances. This is used to evaluate
/// retrieval quality (e.g. Recall@K) and to build N-pair batches.
pub fn siamese_distances(
  embeddings : Array[Array[Float]],
) -> Array[Float] {
  let n = embeddings.length()
  let out : Array[Float] = Array::make(n * n, 0.0F)
  for i in 0.. Int {
  embedding_net_num_params(s.net)
}