// lstm_cell.mbt — Vanilla LSTM cell forward + backward (v0.29.0).
//
// Single LSTM cell step. Inputs at time t:
//   - x_t   : input,  length d_x
//   - h_{t-1} : hidden, length d_h
//   - c_{t-1} : cell state, length d_h
//
// Vanilla LSTM (no peepholes):
//   f_t = σ(W_f · [h_{t-1}; x_t] + b_f)         forget gate
//   i_t = σ(W_i · [h_{t-1}; x_t] + b_i)         input gate
//   C̃_t = tanh(W_C · [h_{t-1}; x_t] + b_C)     candidate cell
//   C_t = f_t ⊙ C_{t-1} + i_t ⊙ C̃_t            cell state
//   o_t = σ(W_o · [h_{t-1}; x_t] + b_o)         output gate
//   h_t = o_t ⊙ tanh(C_t)                      hidden output
//
// This module provides:
//   - `LstmCellParam`  — (w_f, w_i, w_c, w_o, b_f, b_i, b_c, b_o)
//   - `LstmCellCache`  — intermediates for BPTT
//   - `LstmCellGrad`   — parameter gradients (accumulated over time)
//   - `lstm_cell_forward`  — single step forward
//   - `lstm_cell_backward` — single step backward
//   - `lstm_cell_sgd_step` — apply accumulated gradients (SGD)
//
// Note: struct fields MUST be lowercase (MoonBit requirement), so we
// use `w_f` instead of `W_f`. The mathematical notation `W_f` is
// preserved in the doc comments.

///|
/// LSTM cell parameter bundle. Each `w_X` is d_h × (d_h + d_x); each
/// `b_X` has length d_h.
pub struct LstmCellParam {
  d_x : Int
  d_h : Int
  w_f : Array[Array[Float]]
  w_i : Array[Array[Float]]
  w_c : Array[Array[Float]]
  w_o : Array[Array[Float]]
  b_f : Array[Float]
  b_i : Array[Float]
  b_c : Array[Float]
  b_o : Array[Float]
}

///|
/// Build a fresh LSTM cell parameter with Xavier-normal init for
/// the four weight matrices and zero biases. seed controls RNG.
pub fn LstmCellParam::new(
  d_x : Int,
  d_h : Int,
  seed : UInt64,
) -> LstmCellParam {
  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 wf = xavier_normal(d_h, in_dim, std, rng)
  let wi = xavier_normal(d_h, in_dim, std, rng)
  let wc = xavier_normal(d_h, in_dim, std, rng)
  let wo = xavier_normal(d_h, in_dim, std, rng)
  let bf : Array[Float] = Array::make(d_h, 0.0F)
  let bi : Array[Float] = Array::make(d_h, 0.0F)
  let bc : Array[Float] = Array::make(d_h, 0.0F)
  let bo : Array[Float] = Array::make(d_h, 0.0F)
  { d_x, d_h, w_f: wf, w_i: wi, w_c: wc, w_o: wo, b_f: bf, b_i: bi, b_c: bc, b_o: bo }
}

///|
/// Cache of intermediate values from a single LSTM forward step.
pub struct LstmCellCache {
  x : Array[Float]
  h_prev : Array[Float]
  c_prev : Array[Float]
  f : Array[Float]
  ig : Array[Float]
  c_tilde : Array[Float]
  o : Array[Float]
  c_t : Array[Float]
  tanh_c_t : Array[Float]
  h_t : Array[Float]
}

///|
/// Accumulated parameter gradients from a sequence of backward
/// steps. Initialise all entries to zero and accumulate via
/// `lstm_cell_backward` + `lstm_cell_sgd_step`.
pub struct LstmCellGrad {
  d_w_f : Array[Array[Float]]
  d_w_i : Array[Array[Float]]
  d_w_c : Array[Array[Float]]
  d_w_o : Array[Array[Float]]
  d_b_f : Array[Float]
  d_b_i : Array[Float]
  d_b_c : Array[Float]
  d_b_o : Array[Float]
}

///|
/// Create a zero-initialised LstmCellGrad matching the param shape.
pub fn LstmCellGrad::zero(param : LstmCellParam) -> LstmCellGrad {
  let d_h = param.d_h
  let in_dim = param.d_h + param.d_x
  let zwf : Array[Array[Float]] = Array::make(d_h, [])
  let zwi : Array[Array[Float]] = Array::make(d_h, [])
  let zwc : Array[Array[Float]] = Array::make(d_h, [])
  let zwo : Array[Array[Float]] = Array::make(d_h, [])
  for i in 0.. Array[Float] {
  let n = a.length()
  let out : Array[Float] = Array::make(n, 0.0F)
  for i in 0.. Array[Array[Float]] {
  let w : Array[Array[Float]] = Array::make(rows, [])
  for i in 0.. Array[Float] {
  let n = x.length()
  let out : Array[Float] = Array::make(n, 0.0F)
  for i in 0.. Array[Float] {
  let n = x.length()
  let out : Array[Float] = Array::make(n, 0.0F)
  for i in 0.. Array[Float] {
  let d_out = w.length()
  let y : Array[Float] = Array::make(d_out, 0.0F)
  for i in 0.. Array[Float] {
  let d_out = w.length()
  let d_in = x.length()
  let d_x : Array[Float] = Array::make(d_in, 0.0F)
  for i in 0.. Array[Float] {
  let out : Array[Float] = Array::make(a.length() + b.length(), 0.0F)
  for i in 0.. (Array[Float], Array[Float]) {
  let first : Array[Float] = Array::make(n_first, 0.0F)
  let second : Array[Float] = Array::make(v.length() - n_first, 0.0F)
  for i in 0.. (Array[Float], LstmCellCache) {
  let d_h = param.d_h
  let cat = vec_concat(h_prev, x)
  let f_pre = matvec(param.w_f, param.b_f, cat)
  let f = sigmoid_vec(f_pre)
  let i_pre = matvec(param.w_i, param.b_i, cat)
  let ig = sigmoid_vec(i_pre)
  let c_pre = matvec(param.w_c, param.b_c, cat)
  let c_tilde = tanh_vec(c_pre)
  let o_pre = matvec(param.w_o, param.b_o, cat)
  let o = sigmoid_vec(o_pre)
  // Cell state.
  let c_t : Array[Float] = Array::make(d_h, 0.0F)
  for i in 0.. (Array[Float], Array[Float], Array[Float]) {
  let d_h = param.d_h
  // d_o = d_h ⊙ tanh(C_t)
  let d_o : Array[Float] = Array::make(d_h, 0.0F)
  for i in 0.. Unit {
  for i in 0..