// chain.mbt — Generic layer Chain / Sequential (v0.19.0).
//
// Tagged-enum dispatch over the project's forward / backward layers.
// The user composes a sequence of layers, calls chain_forward to run
// them in order, and chain_backward to reverse the order. Caches per
// layer are returned alongside the output, and per-layer parameter
// gradients are returned alongside d_input.
//
// Conventions:
//   - Layer enum carries each layer's parameter struct (or unit for
//     stateless layers like ReLU / Flatten / MaxPool2d).
//   - LayerCache enum carries whatever the backward pass needs (input
//     activations, BN stats, maxpool argmax indices, ...).
//   - LayerGrad enum carries per-layer parameter gradients (Empty for
//     layers with no learnable params).
//   - Shape `(n, c, h, w)` is tracked explicitly. After Flatten or
//     Linear, the chain treats the data as `(n, k, 1, 1)` where `k` is
//     the number of features; subsequent Conv2d / MaxPool2d layers
//     will fail (semantically wrong topology).
//
// Not in scope: residual / skip / branching connections. Those are
// v0.20.x material (Function API).

///|
/// Layer enum: each variant holds the parameter struct (or is unit
/// for stateless layers).
pub enum Layer {
  Conv2d(Conv2dParam)
  ReLU
  MaxPool2d(MaxPool2dParam)
  Flatten
  Linear(LinearParam)
  BatchNorm2d(BatchNorm2d)
  LayerNorm(LayerNorm)
}

// ---------------------------------------------------------------------------
// Public constructor helpers (MoonBit's read-only enum variants cannot
// be constructed from another file; callers must use these wrappers).
// ---------------------------------------------------------------------------

///|
pub fn Layer::conv2d(param : Conv2dParam) -> Layer { Layer::Conv2d(param) }

///|
pub fn Layer::relu() -> Layer { Layer::ReLU }

///|
pub fn Layer::max_pool2d(param : MaxPool2dParam) -> Layer { Layer::MaxPool2d(param) }

///|
pub fn Layer::flatten() -> Layer { Layer::Flatten }

///|
pub fn Layer::linear(param : LinearParam) -> Layer { Layer::Linear(param) }

///|
pub fn Layer::batch_norm2d(bn : BatchNorm2d) -> Layer { Layer::BatchNorm2d(bn) }

///|
pub fn Layer::layer_norm(ln : LayerNorm) -> Layer { Layer::LayerNorm(ln) }

///|
/// Per-layer forward cache. Each variant holds the data the backward
/// pass needs (input activations, BN stats, max-pool argmax indices,
/// etc.).
pub enum LayerCache {
  // input + shape (for conv2d_backward)
  Conv2d(Array[Float], Int, Int, Int, Int)
  // input (for relu_backward)
  ReLU(Array[Float])
  // argmax_idx + shape + kernel/stride/pad (for maxpool2d_backward)
  MaxPool2d(Array[Int], Int, Int, Int, Int, Int, Int, Int, Int)
  // no cache
  Flatten
  // input + n (for linear_backward)
  Linear(Array[Float], Int)
  // full BN cache (mean, inv_std, x_centered, ...)
  BatchNorm2d(BatchNormCache)
  // full LN cache
  LayerNorm(LayerNormCache)
}

///|
/// Per-layer parameter gradients. `Empty` for layers without
/// learnable params (ReLU, Flatten, MaxPool2d).
pub enum LayerGrad {
  Empty
  Conv2d(Array[Float], Array[Float]) // d_weight, d_bias
  Linear(Array[Float], Array[Float])
  BatchNorm2d(Array[Float], Array[Float]) // d_gamma, d_beta
  LayerNorm(Array[Float], Array[Float])
  MaxPool2d
}

///|
/// Forward pass over the chain. Returns the final output, one cache
/// per layer (in forward order), and the final `(n, c, h, w)` shape.
///
/// Note: after Flatten or Linear the chain internally treats the data
/// as `(n, k, 1, 1)`. Subsequent Conv2d / MaxPool2d / BN / LN layers
/// will fail at runtime — this is a semantic topology error.
pub fn chain_forward(
  layers : Array[Layer],
  input : Array[Float],
  n : Int,
  c : Int,
  h : Int,
  w : Int,
) -> (Array[Float], Array[LayerCache], Int, Int, Int, Int) {
  let caches : Array[LayerCache] = []
  let mut data = input
  let cur_n = n
  let mut cur_c = c
  let mut cur_h = h
  let mut cur_w = w
  for i in 0.. {
        caches.push(LayerCache::Conv2d(data, cur_n, cur_c, cur_h, cur_w))
        let out = conv2d_forward(data, cur_n, cur_c, cur_h, cur_w, param)
        data = out
        // Compute new spatial dims: ho = (h + 2*pad - kh) / stride + 1
        let ho = (cur_h + 2 * param.pad - param.kh) / param.stride + 1
        let wo = (cur_w + 2 * param.pad - param.kw) / param.stride + 1
        cur_c = param.c_out
        cur_h = ho
        cur_w = wo
      }
      ReLU => {
        caches.push(LayerCache::ReLU(data))
        data = relu_forward(data)
        // shape unchanged
      }
      MaxPool2d(param) => {
        let (out, argmax) = maxpool2d_forward_with_idx(
          data, cur_n, cur_c, cur_h, cur_w, param,
        )
        caches.push(
          LayerCache::MaxPool2d(
            argmax, cur_n, cur_c, cur_h, cur_w, param.kh, param.kw,
            param.stride, param.pad,
          ),
        )
        data = out
        let ho = (cur_h + 2 * param.pad - param.kh) / param.stride + 1
        let wo = (cur_w + 2 * param.pad - param.kw) / param.stride + 1
        cur_h = ho
        cur_w = wo
      }
      Flatten => {
        caches.push(LayerCache::Flatten)
        // Treat the current (c, h, w) plane as a flat feature vector
        // per sample: feature dim = c * h * w.
        let k = cur_c * cur_h * cur_w
        cur_c = k
        cur_h = 1
        cur_w = 1
        // data is already in the right memory order (NCHW with c*h*w
        // contiguous per sample) so we don't need to copy.
      }
      Linear(param) => {
        caches.push(LayerCache::Linear(data, cur_n))
        let out = linear_forward(data, cur_n, param)
        data = out
        cur_c = param.out_features
        cur_h = 1
        cur_w = 1
      }
      BatchNorm2d(bn) => {
        let (out, cache) = batch_norm2d_forward(data, cur_n, cur_c, cur_h, cur_w, bn)
        caches.push(LayerCache::BatchNorm2d(cache))
        data = out
      }
      LayerNorm(ln) => {
        let (out, cache) = layer_norm_forward(data, cur_n, cur_c, cur_h, cur_w, ln)
        caches.push(LayerCache::LayerNorm(cache))
        data = out
      }
    }
  }
  (data, caches, cur_n, cur_c, cur_h, cur_w)
}

///|
/// Backward pass. Iterates the chain in reverse, threading `d_output`
/// through each layer's backward and emitting per-layer parameter
/// gradients. Returns `(d_input, grads)` where `grads[i]` corresponds
/// to `layers[i]`.
pub fn chain_backward(
  layers : Array[Layer],
  caches : Array[LayerCache],
  d_output : Array[Float],
  n : Int,
  c : Int,
  h : Int,
  w : Int,
) -> (Array[Float], Array[LayerGrad]) {
  let grads : Array[LayerGrad] = Array::make(layers.length(), LayerGrad::Empty)
  let mut d = d_output
  let mut cur_n = n
  let mut cur_c = c
  let mut cur_h = h
  let mut cur_w = w
  // Walk the chain in reverse.
  for i in 0.. {
        d = relu_backward(input, d)
        // shape unchanged
      }
      (Flatten, LayerCache::Flatten) => {
        // d has shape (cur_n, cur_c, 1, 1) where cur_c = c*h*w of the
        // previous (pre-Flatten) tensor. We need to interpret it as
        // (cur_n, prev_c, prev_h, prev_w). The data is already in the
        // right memory order (NCHW row-major with feature dim
        // contiguous), so the only change is logical shape.
        // We can't restore prev_c/h/w here without tracking it, so
        // we leave the data unchanged and pass it forward. The next
        // (earlier) layer's backward receives this flat-shape array
        // but treats it using the shape recorded in its own cache.
        let _ = cur_n
      }
      (MaxPool2d(param), LayerCache::MaxPool2d(argmax, n0, c0, h0, w0, kh, kw, st, pad)) => {
        // The d here is the upstream d, which has shape (n, c, h_out, w_out).
        let d_in = maxpool2d_backward(d, argmax, n0, c0, h0, w0, param)
        // d_in has shape (n, c, h0, w0). Update shape.
        d = d_in
        cur_n = n0
        cur_c = c0
        cur_h = h0
        cur_w = w0
        let _ = kh
        let _ = kw
        let _ = st
        let _ = pad
        grads[idx] = LayerGrad::MaxPool2d
      }
      // ---- Layers with parameters ----
      (Conv2d(param), LayerCache::Conv2d(input, n0, c0, h0, w0)) => {
        let (d_in, d_w, d_b) = conv2d_backward(input, d, n0, c0, h0, w0, param)
        d = d_in
        cur_n = n0
        cur_c = c0
        cur_h = h0
        cur_w = w0
        grads[idx] = LayerGrad::Conv2d(d_w, d_b)
      }
      (Linear(param), LayerCache::Linear(input, n0)) => {
        let (d_in, d_w, d_b) = linear_backward(input, d, n0, param)
        d = d_in
        // d_in has shape (n0, in_features). We don't know in_features
        // here directly — it's the previous layer's out_features or
        // the post-Flatten feature dim. Just mark cur_h = cur_w = 1.
        cur_n = n0
        cur_h = 1
        cur_w = 1
        grads[idx] = LayerGrad::Linear(d_w, d_b)
      }
      (BatchNorm2d(bn), LayerCache::BatchNorm2d(bn_cache)) => {
        let (d_in, d_g, d_b) = batch_norm2d_backward(d, bn_cache, bn)
        d = d_in
        grads[idx] = LayerGrad::BatchNorm2d(d_g, d_b)
      }
      (LayerNorm(ln), LayerCache::LayerNorm(ln_cache)) => {
        let (d_in, d_g, d_b) = layer_norm_backward(d, ln_cache, ln)
        d = d_in
        grads[idx] = LayerGrad::LayerNorm(d_g, d_b)
      }
      // Mismatched cache / layer combos — propagate d unchanged and
      // leave the grad as Empty. (This branch is unreachable when the
      // user passes a well-formed chain.)
      _ => ()
    }
  }
  (d, grads)
}