// wgan_trainer.mbt -- WGAN-GP trainer (v0.128.0).
//
// WGAN-GP alternates two updates per iteration:
//
//   1. n_critic critic steps (the paper uses n_critic = 5):
//        sample z ~ N(0, I)
//        fake = G(z)
//        L_critic = E[f_w(real)] - E[f_w(fake)] + lambda * L_GP
//        descend on L_critic
//   2. one generator step:
//        sample z ~ N(0, I)
//        L_G = -E[f_w(G(z))]
//        descend on L_G
//
// The extra critic steps keep the critic close to optimal before each
// generator update, which is what makes the Wasserstein objective
// well-behaved in practice.
//
// Scope of v0.128.0:
//   - WGAN struct: generator + critic + latent_dim.
//   - wgan_critic_step: one critic round (Wasserstein + GP loss).
//   - wgan_generator_step: one generator round.
//   - WGANTrainer: n_critic loop over a mini-batch.
//   - wgan_eval: held-out Wasserstein score gap + mean gradient norm.
//   - wgan_num_params: total trainable scalars.
//
// 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: Gulrajani et al. 2018 (WGAN-GP).

///|
/// WGAN: composite of a generator and a Wasserstein critic.
pub struct WGAN {
  g : DCGANGenerator
  c : WCritic
  latent_dim : Int
}

///|
/// Build a fresh WGAN.
pub fn WGAN::new(g : DCGANGenerator, c : WCritic) -> WGAN {
  { g, c, latent_dim: g.latent_dim }
}

///|
/// One critic round: sample z, generate a fake image, and evaluate the
/// Wasserstein + gradient-penalty objective. Returns
/// (w_loss, gp_value, fake_image).
pub fn wgan_critic_step(
  wgan : WGAN,
  real_image : Array[Float],
  lambda_gp : Float,
  fd_eps : Float,
  rng : Xoshiro,
) -> (Float, Float, Array[Float]) {
  let z = dcgan_sample_latent(wgan.latent_dim, rng)
  let fake = dcgan_generator_forward(wgan.g, z)
  let w_loss = w_critic_loss(wgan.c, real_image, fake)
  let eps = gp_sample_eps(rng)
  let x_hat = gp_interpolate(real_image, fake, eps)
  let grad = gp_input_gradient(wgan.c, x_hat, fd_eps)
  let norm = gp_grad_norm(grad)
  let diff = norm - 1.0F
  let gp = diff * diff
  let _ = lambda_gp
  (w_loss, gp, fake)
}

///|
/// One generator round: sample z, generate a fake image, and evaluate
/// the generator objective L_G = -E[f_w(G(z))]. Returns
/// (g_loss, fake_image).
pub fn wgan_generator_step(
  wgan : WGAN,
  rng : Xoshiro,
) -> (Float, Array[Float]) {
  let z = dcgan_sample_latent(wgan.latent_dim, rng)
  let fake = dcgan_generator_forward(wgan.g, z)
  let score = w_critic_score(wgan.c, fake)
  (-score, fake)
}

///|
/// WGANTrainer: wraps a WGAN with training hyperparameters.
pub struct WGANTrainer {
  wgan : WGAN
  n_critic : Int
  lambda_gp : Float
  fd_eps : Float
  batch_size : Int
}

///|
/// Build a fresh WGANTrainer. `n_critic` is the number of critic steps
/// per generator step (the paper uses 5); `lambda_gp` is the
/// gradient-penalty coefficient (the paper uses 10).
pub fn WGANTrainer::new(
  wgan : WGAN,
  n_critic : Int,
  lambda_gp : Float,
  fd_eps : Float,
  batch_size : Int,
) -> WGANTrainer {
  { wgan, n_critic, lambda_gp, fd_eps, batch_size }
}

///|
/// One full training iteration on a mini-batch: n_critic critic rounds
/// followed by one generator round. Returns the mean critic loss
/// (Wasserstein + lambda * GP), the mean gradient-penalty value, and
/// the generator loss.
pub fn wgan_train_step(
  trainer : WGANTrainer,
  real_batch : Array[Array[Float]],
  rng : Xoshiro,
) -> (Float, Float, Float) {
  let m = real_batch.length()
  let mut w_total = 0.0F
  let mut gp_total = 0.0F
  let mut g_total = 0.0F
  let n_iter = trainer.n_critic
  for i in 0.. (Float, Float, Float) {
  let rng = Xoshiro::new(0UL)
  let m = held_out.length()
  let mut real_sum = 0.0F
  let mut fake_sum = 0.0F
  let mut norm_sum = 0.0F
  for i in 0.. Int {
  let critic_params = dcgan_discriminator_num_params(w_critic_disc(wgan.c))
  dcgan_generator_num_params(wgan.g) + critic_params
}