// batch_norm2d.mbt — 2D Batch Normalisation for NCHW tensors (v0.17.0).
//
// Reference: Ioffe & Szegedy, "Batch Normalization: Accelerating Deep
// Network Training by Reducing Internal Covariate Shift", ICML 2015.
//
// Forward (training mode, per channel c):
//   mean_c     = sum over (n,h,w) of x[n,c,h,w] / (n*h*w)
//   var_c      = sum over (n,h,w) of (x[n,c,h,w] - mean_c)^2 / (n*h*w)
//   inv_std_c  = 1 / sqrt(var_c + eps)
//   y[n,c,h,w] = gamma[c] * (x[n,c,h,w] - mean_c) * inv_std_c + beta[c]
//
// Backward: returns d_input / d_gamma / d_beta given d_output. Uses the
// standard "save mean + inv_std + x_centered" cache to avoid recomputation.
//
// Layout: NCHW row-major flat `Array[Float]`.
//   Input:    [n, c, h, w]    length = n * c  * h  * w
//   Output:   [n, c, h, w]    length = n * c  * h  * w
//   Gamma:    [c]             length = c
//   Beta:     [c]             length = c
//   Running mean / var: [c]
//
// Conventions match the project's bit-exact Float32 contract. We do not
// use Bessel's correction (population variance, divisor = N*H*W) since
// that matches PyTorch's `track_running_stats=False` training behaviour.

///|
/// BatchNorm2d parameter container.
pub struct BatchNorm2d {
  // Learnable affine parameters.
  gamma : Array[Float]
  beta : Array[Float]
  // Running statistics (for inference); updated with exponential moving
  // average during training. Empty if running stats are not tracked.
  running_mean : Array[Float]
  running_var : Array[Float]
  // BN hyper-parameters.
  momentum : Float  // EMA coefficient for running stats; 0 disables updates.
  eps : Float
  // Training mode flag — when false, forward uses running_mean / running_var
  // instead of batch statistics.
  mut training : Bool
}

///|
/// Build a BatchNorm2d for `c` channels. `gamma` is initialised to all-1,
/// `beta` to all-0; running stats to all-0. `momentum` defaults to 0.1,
/// `eps` to 1e-5.
pub fn BatchNorm2d::new(c : Int, momentum? : Float = 0.1, eps? : Float = 0.00001) -> BatchNorm2d {
  {
    gamma: Array::make(c, 1.0F),
    beta: Array::make(c, 0.0F),
    running_mean: Array::make(c, 0.0F),
    running_var: Array::make(c, 0.0F),
    momentum,
    eps,
    training: true,
  }
}

///|
/// Convenience constructor taking pre-built gamma / beta arrays (does NOT
/// copy them).
pub fn BatchNorm2d::with_gamma_beta(
  gamma : Array[Float],
  beta : Array[Float],
  momentum? : Float = 0.1,
  eps? : Float = 0.00001,
) -> BatchNorm2d {
  let c = gamma.length()
  {
    gamma,
    beta,
    running_mean: Array::make(c, 0.0F),
    running_var: Array::make(c, 0.0F),
    momentum,
    eps,
    training: true,
  }
}

///|
/// Set training / inference mode.
pub fn BatchNorm2d::set_training(self : BatchNorm2d, training : Bool) -> Unit {
  self.training = training
}

///|
/// Cache returned by `batch_norm2d_forward` and consumed by
/// `batch_norm2d_backward`. Holds everything the backward pass needs
/// without recomputing the forward.
pub struct BatchNormCache {
  // Per-channel statistics computed from the current batch.
  mean : Array[Float]
  inv_std : Array[Float]
  // (n, c, h, w) — input shifted by the per-channel mean.
  x_centered : Array[Float]
  // n * h * w per channel; cached to avoid re-multiplying in backward.
  nhw : Int
  // Shape, for backward.
  n : Int
  c : Int
  h : Int
  w : Int
}

///|
/// Forward pass. `bn` may be in training mode (uses batch stats) or
/// inference mode (uses running_mean / running_var). When `training=true`
/// and `bn.momentum > 0`, the running stats are updated in-place.
pub fn batch_norm2d_forward(
  input : Array[Float],
  n : Int,
  c : Int,
  h : Int,
  w : Int,
  bn : BatchNorm2d,
) -> (Array[Float], BatchNormCache) {
  let nhw = n * h * w
  let chw = c * h * w
  let total = n * chw
  // Allocate output and cache buffers up front.
  let output : Array[Float] = Array::make(total, 0.0F)
  let x_centered : Array[Float] = Array::make(total, 0.0F)
  let mean : Array[Float] = Array::make(c, 0.0F)
  let inv_std : Array[Float] = Array::make(c, 0.0F)
  // Pick which mean / var to use based on training mode.
  let use_running = not(bn.training)
  let use_mean : Array[Float] = if use_running { bn.running_mean } else { mean }
  let use_var : Array[Float] = if use_running {
    bn.running_var
  } else {
    // We'll fill `mean` and compute `inv_std` from `var` below.
    mean // placeholder; inv_std computed separately
  }
  // ---- Per-channel batch statistics (always computed for the cache) ----
  // We always compute batch mean / var so the cache is complete (in
  // inference mode we still want the cache to record what we used).
  for ci in 0.. (Array[Float], Array[Float], Array[Float]) {
  let n = cache.n
  let c = cache.c
  let h = cache.h
  let w = cache.w
  let nhw = cache.nhw
  let chw = c * h * w
  let total = n * chw
  let d_input : Array[Float] = Array::make(total, 0.0F)
  let d_gamma : Array[Float] = Array::make(c, 0.0F)
  let d_beta : Array[Float] = Array::make(c, 0.0F)
  let nhw_f = Float::from_int(nhw)
  for ci in 0..