// layer_norm.mbt — Layer Normalisation for NCHW tensors (v0.17.1).
//
// Reference: Ba, Kiros, Hinton, "Layer Normalization", arXiv:1607.06450
// (2016). Normalises per-sample across the feature axes (C, H, W).
//
// Forward (training, per sample n):
//   mean_n     = sum over (c, h, w) of x[n,c,h,w] / (c*h*w)
//   var_n      = sum over (c, h, w) of (x[n,c,h,w] - mean_n)^2 / (c*h*w)
//   inv_std_n  = 1 / sqrt(var_n + eps)
//   y[n,c,h,w] = gamma[c,h,w] * (x[n,c,h,w] - mean_n) * inv_std_n
//                + beta[c,h,w]
//
// Backward: closed-form per-sample using saved mean + inv_std + x_centered.
//
// 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 * h * w]     length = c * h * w (per-feature scale)
//   Beta:   [c * h * w]     length = c * h * w (per-feature shift)
//
// No running stats — LayerNorm is deterministic (unlike BN which tracks
// running_mean / running_var for inference).

///|
/// LayerNorm parameter container. Affine params have one element per
/// feature position (length = c * h * w).
pub struct LayerNorm {
  gamma : Array[Float]
  beta : Array[Float]
  eps : Float
  c : Int
  h : Int
  w : Int
}

///|
/// Build a LayerNorm for the given feature shape (c, h, w). `gamma` is
/// initialised to all-1, `beta` to all-0.
pub fn LayerNorm::new(c : Int, h : Int, w : Int, eps? : Float = 0.00001) -> LayerNorm {
  let chw = c * h * w
  {
    gamma: Array::make(chw, 1.0F),
    beta: Array::make(chw, 0.0F),
    eps,
    c,
    h,
    w,
  }
}

///|
/// Convenience constructor with explicit gamma / beta arrays (no copy).
pub fn LayerNorm::with_gamma_beta(
  gamma : Array[Float],
  beta : Array[Float],
  c : Int,
  h : Int,
  w : Int,
  eps? : Float = 0.00001,
) -> LayerNorm {
  { gamma, beta, eps, c, h, w }
}

///|
/// Cache returned by `layer_norm_forward` and consumed by
/// `layer_norm_backward`.
pub struct LayerNormCache {
  // Per-sample statistics.
  mean : Array[Float]
  inv_std : Array[Float]
  // (n, c, h, w) — input shifted by per-sample mean.
  x_centered : Array[Float]
  // Shape for backward.
  n : Int
  c : Int
  h : Int
  w : Int
  // chw per sample (denominator for backward).
  chw : Int
}

///|
/// Forward pass.
pub fn layer_norm_forward(
  input : Array[Float],
  n : Int,
  c : Int,
  h : Int,
  w : Int,
  ln : LayerNorm,
) -> (Array[Float], LayerNormCache) {
  let chw = c * h * w
  let total = n * chw
  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(n, 0.0F)
  let inv_std : Array[Float] = Array::make(n, 0.0F)
  let chw_f = Float::from_int(chw)
  for ni 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 chw = cache.chw
  let total = n * chw
  let d_input : Array[Float] = Array::make(total, 0.0F)
  let d_gamma : Array[Float] = Array::make(chw, 0.0F)
  let d_beta : Array[Float] = Array::make(chw, 0.0F)
  let chw_f = Float::from_int(chw)
  let inv_chw = 1.0F / chw_f
  for ni in 0..