// wgan_critic.mbt -- Wasserstein critic for WGAN (v0.126.0).
//
// WGAN (Arjovsky et al. 2017 "Wasserstein GAN") replaces the
// discriminator's binary classification loss with a Wasserstein
// distance estimate. Rather than classifying real/fake, the critic
// learns a 1-Lipschitz function f_w that scores real samples high and
// fake samples low, and the loss becomes:
//
//   L_critic = E_{x~data}[f_w(x)] - E_{z~N(0,I)}[f_w(G(z))]
//
// The critic is *minimised* (unlike a discriminator, which is
// maximised). The generator then minimises -E[f_w(G(z))], i.e. pushes
// its samples toward higher critic scores.
//
// The 1-Lipschitz constraint is enforced either by weight clipping
// (WGAN) or by a gradient penalty (WGAN-GP, Gulrajani et al. 2018);
// both are shipped in this project.
//
// Scope of v0.126.0:
//   - WCritic struct: wraps a DCGANDiscriminator as a critic.
//   - w_critic_score: raw critic output f_w(x) for one image.
//   - w_critic_score_batch: critic outputs for a batch.
//   - w_critic_loss: E[real] - E[fake] on one real/fake pair.
//   - w_critic_clip_weights: in-place Lipschitz weight clipping.
//
// Reference: Arjovsky et al. 2017; Gulrajani et al. 2018.

///|
/// WCritic: a Wasserstein critic. Structurally identical to a
/// DCGANDiscriminator (same conv stack) but trained with a different
/// objective, so it is modelled as a thin wrapper rather than a new
/// conv stack.
pub struct WCritic {
  d : DCGANDiscriminator
}

///|
/// Build a fresh WCritic from an existing discriminator-shaped stack.
pub fn WCritic::new(d : DCGANDiscriminator) -> WCritic {
  { d, }
}

///|
/// Raw critic score f_w(x). In WGAN this is a real-valued score (not a
/// probability), so we bypass the sigmoid that the discriminator uses.
pub fn w_critic_score(c : WCritic, image : Array[Float]) -> Float {
  dcgan_discriminator_forward(c.d, image)
}

///|
/// Critic scores for a batch of images.
pub fn w_critic_score_batch(
  c : WCritic,
  batch : Array[Array[Float]],
) -> Array[Float] {
  let m = batch.length()
  let out : Array[Float] = Array::make(m, 0.0F)
  for i in 0.. Float {
  w_critic_score(c, real_image) - w_critic_score(c, fake_image)
}

///|
/// In-place weight clipping to enforce the 1-Lipschitz constraint of
/// the critic (the original WGAN approach; WGAN-GP uses a gradient
/// penalty instead). Clips every weight and bias to [-clip, clip].
pub fn w_critic_clip_weights(c : WCritic, clip : Float) -> Unit {
  let lo = -clip
  let hi = clip
  // Conv 1.
  for o in 0.. hi { hi } else { v } }
    }
  }
  for i in 0.. hi { hi } else { v } }
  }
  // Conv 2.
  for o in 0.. hi { hi } else { v } }
    }
  }
  for i in 0.. hi { hi } else { v } }
  }
  // Conv 3.
  for o in 0.. hi { hi } else { v } }
    }
  }
  for i in 0.. hi { hi } else { v } }
  }
  // Classifier.
  for k in 0.. hi { hi } else { v } }
  }
  c.d.fc_b = if c.d.fc_b < lo {
    lo
  } else {
    if c.d.fc_b > hi { hi } else { c.d.fc_b }
  }
}

///|
/// Access the underlying discriminator (for sampling / param counting).
pub fn w_critic_disc(c : WCritic) -> DCGANDiscriminator {
  c.d
}