// siamese_trainer.mbt -- Siamese training loop (v0.132.0).
//
// Scope of v0.132.0:
//   - SiameseBatch: (anchors, positives, negatives) triplet batch.
//   - build_siamese_batch: construct a batch by sampling positives of
//     the same class and negatives of a different class.
//   - siamese_triplet_step: forward the batch, compute the mean triplet
//     loss, and return it.
//   - SiameseTrainer struct + one training iteration over a batch.
//   - recall_at_k: retrieval metric over a query/gallery/label split.
//
// Parameter updates are deferred to a follow-up batch, consistent with
// the forward-only pattern used across this project.
//
// Reference: FaceNet (Schroff et al. 2015); N-pair (Soh et al. 2016).

///|
/// SiameseBatch: a set of (anchor, positive, negative) triplets.
pub struct SiameseBatch {
  anchors : Array[Array[Float]]
  positives : Array[Array[Float]]
  negatives : Array[Array[Float]]
  size : Int
}

///|
/// Build a triplet batch. For each anchor at index `i`:
///   - the positive is another sample with the same label,
///   - the negative is a sample with a different label.
///
/// Sampling is deterministic given the RNG. Anchors whose class has
/// fewer than two members are skipped (they have no valid positive).
pub fn build_siamese_batch(
  samples : Array[Array[Float]],
  labels : Array[Int],
  batch_size : Int,
  rng : Xoshiro,
) -> SiameseBatch {
  let n = samples.length()
  let anchors : Array[Array[Float]] = Array::make(batch_size, Array::make(0, 0.0F))
  let positives : Array[Array[Float]] = Array::make(batch_size, Array::make(0, 0.0F))
  let negatives : Array[Array[Float]] = Array::make(batch_size, Array::make(0, 0.0F))
  let n_u64 = n.to_uint64()
  let mut filled = 0
  let mut attempts = 0
  let max_attempts = batch_size * 20
  while filled < batch_size && attempts < max_attempts {
    attempts = attempts + 1
    if n == 0 {
      break
    }
    let ai = (next_u64(rng) % n_u64).to_int()
    let label = labels[ai]
    // Find a positive with the same label.
    let pi = (next_u64(rng) % n_u64).to_int()
    if pi == ai || labels[pi] != label {
      continue
    }
    // Find a negative with a different label.
    let ni = (next_u64(rng) % n_u64).to_int()
    if labels[ni] == label {
      continue
    }
    anchors[filled] = samples[ai]
    positives[filled] = samples[pi]
    negatives[filled] = samples[ni]
    filled = filled + 1
  }
  { anchors, positives, negatives, size: filled }
}

///|
/// Forward a triplet batch through the shared weights and return the
/// mean triplet loss plus the mean positive / negative squared
/// distances (useful for monitoring whether the embedding is actually
/// separating the classes).
pub fn siamese_triplet_step(
  s : SiameseNet,
  batch : SiameseBatch,
) -> (Float, Float, Float) {
  if batch.size == 0 {
    return (0.0F, 0.0F, 0.0F)
  }
  let mut loss_sum = 0.0F
  let mut pos_sum = 0.0F
  let mut neg_sum = 0.0F
  for i in 0.. SiameseTrainer {
  { s, lr, batch_size }
}

///|
/// One training iteration: build a batch from the pool and evaluate the
/// triplet objective. Returns (loss, mean_pos_d2, mean_neg_d2).
pub fn siamese_train_step(
  trainer : SiameseTrainer,
  samples : Array[Array[Float]],
  labels : Array[Int],
  rng : Xoshiro,
) -> (Float, Float, Float) {
  let batch = build_siamese_batch(
    samples, labels, trainer.batch_size, rng,
  )
  let (loss, pd, nd) = siamese_triplet_step(trainer.s, batch)
  let _ = trainer.lr
  (loss, pd, nd)
}

///|
/// Recall@K over an embedding set with known labels. For each query,
/// count how many of the K nearest gallery items share its label. This
/// is the standard retrieval metric for face / image retrieval.
pub fn recall_at_k(
  s : SiameseNet,
  queries : Array[Array[Float]],
  gallery : Array[Array[Float]],
  labels : Array[Int],
  k : Int,
) -> Float {
  let nq = queries.length()
  let ng = gallery.length()
  if nq == 0 || ng == 0 {
    return 0.0F
  }
  let kk = if k > ng { ng } else { k }
  // Pre-embed the gallery.
  let gallery_emb : Array[Array[Float]] = Array::make(
    ng, Array::make(0, 0.0F),
  )
  for j in 0..= 0 {
        if labels[best_idx] == label {
          hits = hits + 1.0F
        }
        // Mark as selected by setting its distance to +inf.
        best_d[best_idx] = 1.0e30F
        total = total + 1.0F
      }
    }
  }
  if total == 0.0F {
    0.0F
  } else {
    hits / total
  }
}