// 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..