// dcgan_trainer.mbt -- DCGAN trainer (v0.124.0).
//
// Scope of v0.124.0:
//   - DCGANTrainer struct (DCGAN + learning rates + batch_size).
//   - dcgan_train_step: one mini-batch adversarial round; returns the
//     mean (d_loss, g_loss) over the batch.
//   - dcgan_eval: held-out D/G losses plus a realism/accuracy metric.
//   - dcgan_sample_grid_stats: simple summary statistics of generated
//     samples (min / max / mean), useful for monitoring collapse.
//   - dcgan_num_params: total trainable scalars in G + D.
//
// Parameter updates (backward through the conv stacks) are deferred to
// a follow-up batch, consistent with the forward-only pattern used
// across this project.
//
// Reference: Radford et al. 2016.

///|
/// DCGANTrainer: pairs a DCGAN with training hyperparameters.
pub struct DCGANTrainer {
  gan : DCGAN
  lr_g : Float
  lr_d : Float
  batch_size : Int
}

///|
/// Build a fresh DCGANTrainer.
pub fn DCGANTrainer::new(
  gan : DCGAN,
  lr_g : Float,
  lr_d : Float,
  batch_size : Int,
) -> DCGANTrainer {
  { gan, lr_g, lr_d, batch_size }
}

///|
/// One mini-batch adversarial training round. Returns the mean
/// discriminator loss and mean generator loss over the batch.
pub fn dcgan_train_step(
  trainer : DCGANTrainer,
  real_batch : Array[Array[Float]],
  rng : Xoshiro,
) -> (Float, Float) {
  let m = real_batch.length()
  let mut d_total = 0.0F
  let mut g_total = 0.0F
  for i in 0.. (Float, Float, Float) {
  let rng = Xoshiro::new(0UL)
  let m = held_out.length()
  let mut d_total = 0.0F
  let mut g_total = 0.0F
  let mut correct = 0.0F
  for i in 0.. 0.5).
    if dcgan_discriminator_prob(trainer.gan.d, held_out[i]) > 0.5F {
      correct = correct + 1.0F
    }
    // Fake should be classified as fake (prob < 0.5).
    if dcgan_discriminator_prob(trainer.gan.d, fake) < 0.5F {
      correct = correct + 1.0F
    }
  }
  let n = Float::from_int(m)
  (d_total / n, g_total / n, correct / (2.0F * n))
}

///|
/// Summary statistics of generated samples: (min, max, mean). Useful
/// for spotting mode collapse (a collapsed generator produces samples
/// with near-zero variance) and saturation (tanh output pinned at
/// +/-1).
pub fn dcgan_sample_grid_stats(
  gan : DCGAN,
  n_samples : Int,
  rng : Xoshiro,
) -> (Float, Float, Float) {
  let batch = dcgan_generate_batch(gan, n_samples, rng)
  let first = batch[0]
  let mut lo = first[0]
  let mut hi = first[0]
  let mut sum = 0.0F
  let mut count = 0.0F
  for s in 0.. hi {
        hi = img[i]
      }
      sum = sum + img[i]
      count = count + 1.0F
    }
  }
  (lo, hi, sum / count)
}

///|
/// Total trainable scalars across the generator and the discriminator.
pub fn dcgan_num_params(gan : DCGAN) -> Int {
  dcgan_generator_num_params(gan.g) + dcgan_discriminator_num_params(gan.d)
}