// ebm_trainer.mbt -- EBM trainer via Denoising Score Matching (v0.120.0).
//
// Scope of v0.120.0:
//   - EBMTrainer struct (EBM + learning rate).
//   - ebm_train_step: one Denoising Score Matching (DSM) update:
//       1. Sample z ~ N(0, I).
//       2. x_tilde = x + sigma * z.
//       3. score = -dE/dx_tilde (FD approximation, via langevin_gradient).
//       4. loss = 0.5 * mean_i (score_i + z_i / sigma)^2.
//     BPTT-driven parameter updates are deferred to a follow-up batch.
//   - ebm_train_batch: mean loss over a mini-batch.
//   - ebm_eval_loss: deterministic held-out DSM loss.
//   - ebm_num_params: learnable scalar count.
//
// Reference: Vincent 2011 "A Connection Between Score Matching and
// Denoising Autoencoders".

///|
/// EBMTrainer: thin wrapper pairing an EBM with a learning rate.
pub struct EBMTrainer {
  ebm : EBM
  lr : Float
  sigma : Float
}

///|
/// Build a fresh EBMTrainer. `sigma` is the noise scale for DSM
/// (typical values: 0.1 to 1.0 depending on data scale).
pub fn EBMTrainer::new(ebm : EBM, lr : Float, sigma : Float) -> EBMTrainer {
  { ebm, lr, sigma }
}

///|
/// One DSM training step. Returns the loss. Energy parameters are NOT
/// updated in this version.
pub fn ebm_train_step(
  trainer : EBMTrainer,
  x : Array[Float],
  rng : Xoshiro,
) -> Float {
  let n = x.length()
  let sigma = trainer.sigma
  // Sample z ~ N(0, I) and build x_tilde = x + sigma * z.
  let z : Array[Float] = Array::make(n, 0.0F)
  let pairs = n / 2
  let mut i = 0
  while i < pairs * 2 {
    let (z1d, z2d) = box_muller(rng)
    z[i] = Float::from_double(z1d)
    if i + 1 < n {
      z[i + 1] = Float::from_double(z2d)
    }
    i = i + 2
  }
  if n % 2 == 1 {
    let (z1d, _) = box_muller(rng)
    z[n - 1] = Float::from_double(z1d)
  }
  let x_tilde : Array[Float] = Array::make(n, 0.0F)
  for k in 0.. Float {
  let m = batch.length()
  let mut total = 0.0F
  for e in 0.. Float {
  let rng = Xoshiro::new(0UL)
  let m = batch.length()
  let mut total = 0.0F
  let n = batch[0].length()
  let sigma = trainer.sigma
  // Draw one batch of z (one per example) up-front so eval is reproducible.
  for e in 0.. Int {
  let net = ebm.net
  let mut n_params = 0
  n_params = n_params + net.w1.length() * net.w1[0].length()
  n_params = n_params + net.b1.length()
  for l in 0..