// gru_cell.mbt — Vanilla GRU cell forward + backward + SGD (v0.30.0).
//
// Single GRU cell step. Inputs at time t:
//   - x_t      : input, length d_x
//   - h_{t-1}  : hidden, length d_h
//
// Vanilla GRU (no reset-after-Fourier tricks, no layer-norm):
//   z_t = σ(W_z · [h_{t-1}; x_t] + b_z)         update gate
//   r_t = σ(W_r · [h_{t-1}; x_t] + b_r)         reset gate
//   s_t = r_t ⊙ h_{t-1}                        reset-applied hidden
//   n_t = tanh(W_n · [s_t; x_t] + b_n)         candidate
//   h_t = (1 - z_t) ⊙ h_{t-1} + z_t ⊙ n_t     hidden output
//
// This module provides:
//   - `GruCellParam`  — (w_z, w_r, w_n, b_z, b_r, b_n)
//   - `GruCellCache`  — intermediates for BPTT
//   - `GruCellGrad`   — parameter gradients (accumulated over time)
//   - `gru_cell_forward`  — single step forward
//   - `gru_cell_backward` — single step backward
//   - `gru_cell_sgd_step` — apply accumulated gradients (SGD)
//
// Like LSTM, struct fields MUST be lowercase (MoonBit requirement),
// so we use `w_z` instead of `W_z`. The mathematical notation `W_z`
// is preserved in the doc comments.

///|
/// GRU cell parameter bundle. Each `w_X` is d_h × (d_h + d_x); each
/// `b_X` has length d_h. W_n also has shape d_h × (d_h + d_x), but
/// its first d_h input slots see `r ⊙ h_prev` (not `h_prev`).
pub struct GruCellParam {
  d_x : Int
  d_h : Int
  w_z : Array[Array[Float]]
  w_r : Array[Array[Float]]
  w_n : Array[Array[Float]]
  b_z : Array[Float]
  b_r : Array[Float]
  b_n : Array[Float]
}

///|
/// Build a fresh GRU cell parameter with Xavier-normal init for the
/// three weight matrices and zero biases. seed controls RNG.
pub fn GruCellParam::new(
  d_x : Int,
  d_h : Int,
  seed : UInt64,
) -> GruCellParam {
  let in_dim = d_h + d_x
  let std = sqrtf(2.0F / Float::from_int(d_h + in_dim))
  let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let wz = xavier_normal(d_h, in_dim, std, rng)
  let wr = xavier_normal(d_h, in_dim, std, rng)
  let wn = xavier_normal(d_h, in_dim, std, rng)
  let bz : Array[Float] = Array::make(d_h, 0.0F)
  let br : Array[Float] = Array::make(d_h, 0.0F)
  let bn : Array[Float] = Array::make(d_h, 0.0F)
  { d_x, d_h, w_z: wz, w_r: wr, w_n: wn, b_z: bz, b_r: br, b_n: bn }
}

///|
/// Cache of intermediate values from a single GRU forward step.
pub struct GruCellCache {
  x : Array[Float]
  h_prev : Array[Float]
  z : Array[Float]
  r : Array[Float]
  s : Array[Float]
  n : Array[Float]
  h_t : Array[Float]
}

///|
/// Accumulated parameter gradients from a sequence of backward
/// steps. Initialise all entries to zero and accumulate via
/// `gru_cell_backward` + `gru_cell_sgd_step`.
pub struct GruCellGrad {
  d_w_z : Array[Array[Float]]
  d_w_r : Array[Array[Float]]
  d_w_n : Array[Array[Float]]
  d_b_z : Array[Float]
  d_b_r : Array[Float]
  d_b_n : Array[Float]
}

///|
/// Create a zero-initialised GruCellGrad matching the param shape.
pub fn GruCellGrad::zero(param : GruCellParam) -> GruCellGrad {
  let d_h = param.d_h
  let in_dim = param.d_h + param.d_x
  let zwz : Array[Array[Float]] = Array::make(d_h, [])
  let zwr : Array[Array[Float]] = Array::make(d_h, [])
  let zwn : Array[Array[Float]] = Array::make(d_h, [])
  for i in 0.. (Array[Float], GruCellCache) {
  let d_h = param.d_h
  let cat = vec_concat(h_prev, x)
  let z_pre = matvec(param.w_z, param.b_z, cat)
  let z = sigmoid_vec(z_pre)
  let r_pre = matvec(param.w_r, param.b_r, cat)
  let r = sigmoid_vec(r_pre)
  // s = r ⊙ h_prev
  let s : Array[Float] = Array::make(d_h, 0.0F)
  for i in 0.. (Array[Float], Array[Float]) {
  let d_h = param.d_h
  let d_x = param.d_x
  let in_dim = d_h + d_x
  // 1. From h_t = (1-z) ⊙ h_prev + z ⊙ n:
  //    d_z  = d_h ⊙ (n - h_prev)
  //    d_n  = d_h ⊙ z
  //    d_h_prev_main = d_h ⊙ (1 - z)
  let d_z : Array[Float] = Array::make(d_h, 0.0F)
  let d_n : Array[Float] = Array::make(d_h, 0.0F)
  let d_h_prev_main : Array[Float] = Array::make(d_h, 0.0F)
  for i in 0.. Unit {
  for i in 0..