// ddpm_trainer.mbt -- DDPM trainer (v0.116.0).
//
// Scope of v0.116.0:
//   - DDPMTrainer struct (wraps a DDPM + learning rate).
//   - ddpm_train_step: sample t, q_sample, compute MSE loss, return loss.
//     (BPTT-driven SGD through the ScoreNetwork MLP is deferred to a
//     follow-up batch — the same way v0.107/v0.108 only ship the
//     forward pass + loss for VAE/IWAE.)
//   - ddpm_train_batch: mean loss over a mini-batch.
//   - ddpm_eval_loss: held-out loss (no dropout / no noise bookkeeping
//     beyond what DDPM already does).
//
// Reference: Ho et al. 2020 (training objective).

///|
/// DDPMTrainer: thin wrapper that pairs a DDPM with a learning rate
/// (used by future BPTT-driven variants).
pub struct DDPMTrainer {
  ddpm : DDPM
  lr : Float
}

///|
/// Build a fresh DDPMTrainer.
pub fn DDPMTrainer::new(ddpm : DDPM, lr : Float) -> DDPMTrainer {
  { ddpm, lr }
}

///|
/// One training step:
///   1. Sample a random timestep t in [0, T).
///   2. Run `q_sample_pair(x0, t, ...)` to obtain (x_t, noise).
///   3. Compute the MSE loss between predicted and true noise.
/// Returns the loss. The score network parameters are NOT updated in
/// this version (BPTT is deferred).
pub fn ddpm_train_step(
  trainer : DDPMTrainer,
  x0 : Array[Float],
  rng : Xoshiro,
) -> Float {
  let num_T = trainer.ddpm.schedule.num_steps
  // Sample t uniformly in [0, T-1].
  let t = (next_u64(rng) % num_T.to_uint64()).to_int()
  let (xt, noise) = q_sample_pair(x0, t, trainer.ddpm.schedule, rng)
  // Predict noise from x_t.
  let pred = score_network_predict(trainer.ddpm.net, xt, t)
  let n = trainer.ddpm.input_dim
  let mut sum = 0.0F
  for i in 0.. Float {
  let mut total = 0.0F
  let m = batch.length()
  for e in 0.. Float {
  let rng = Xoshiro::new(0UL)
  let mut total = 0.0F
  let m = batch.length()
  for e in 0.. Int {
  let net = ddpm.net
  let mut total = 0
  total = total + net.w_in.length() * net.w_in[0].length()
  total = total + net.b_in.length()
  for l in 0..