// gru_forecaster.mbt — GRU-based time-series forecaster with full BPTT
// + SGD training (v0.87.0).
//
// Architecture (parallel to v0.86.0 LSTMForecaster but with GRU cell):
//
//   input_t (length input_dim)
//       ↓ Linear_in (input_dim → hidden_dim)
//   x_proj_t (length hidden_dim)
//       ↓ GruCell (over time)
//   h_t (length hidden_dim)
//       ↓ Linear_out (hidden_dim → output_dim)
//   y_hat_t (length output_dim)            — point forecast at step t+1
//
// Loss: MSE between y_hat_t and target_{t+1}. Training: full BPTT
// through the GRU cell via `gru_cell_backward` plus analytic gradients
// through Linear_in / Linear_out. Single SGD step on all weights per
// `train_step`.
//
// Reference: Cho et al. 2014 "Learning Phrase Representations using
// RNN Encoder-Decoder for Statistical Machine Translation".

///|
/// GRU time-series forecaster. Projects input → hidden via Linear_in,
/// runs a GruCell over time, projects hidden → output via Linear_out.
pub struct GRUForecaster {
  input_dim : Int
  hidden_dim : Int
  output_dim : Int
  // Linear_in: (hidden_dim × input_dim) + bias of length hidden_dim
  in_w : Array[Array[Float]]
  in_b : Array[Float]
  // GRU cell
  gru : GruCellParam
  // Linear_out: (output_dim × hidden_dim) + bias of length output_dim
  out_w : Array[Array[Float]]
  out_b : Array[Float]
}

///|
/// Build fresh GRUForecaster. Weights init via xavier_normal scaled
/// by sqrt(2/fan_in) (matches v0.86.0 LSTMForecaster init style).
pub fn GRUForecaster::new(
  input_dim : Int,
  hidden_dim : Int,
  output_dim : Int,
  seed : UInt64,
) -> GRUForecaster {
  let rng1 = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let std_in = sqrtf(2.0F / Float::from_int(input_dim))
  let in_w = xavier_normal(hidden_dim, input_dim, std_in, rng1)
  let in_b : Array[Float] = Array::make(hidden_dim, 0.0F)
  let gru = GruCellParam::new(input_dim, hidden_dim, seed + 4UL)
  let rng2 = Xoshiro::from_state(
    seed + 8UL, seed + 9UL, seed + 10UL, seed + 11UL,
  )
  let std_out = sqrtf(2.0F / Float::from_int(hidden_dim))
  let out_w = xavier_normal(output_dim, hidden_dim, std_out, rng2)
  let out_b : Array[Float] = Array::make(output_dim, 0.0F)
  {
    input_dim,
    hidden_dim,
    output_dim,
    in_w,
    in_b,
    gru,
    out_w,
    out_b,
  }
}

///|
/// Per-step forward. Returns `(y_hat, h_t, cache)`. The GRU cell has
/// only a hidden state (no separate cell state).
fn gru_forecaster_step_with_cache(
  model : GRUForecaster,
  input_t : Array[Float],
  h_prev : Array[Float],
) -> (Array[Float], Array[Float], GruCellCache) {
  // x_proj = in_w · input + in_b
  let x_proj = matvec(model.in_w, model.in_b, input_t)
  // h_t, cache = GRU_cell(x_proj, h_prev)
  let (h_t, cache) = gru_cell_forward(x_proj, h_prev, model.gru)
  // y_hat = out_w · h_t + out_b
  let y_hat = matvec(model.out_w, model.out_b, h_t)
  (y_hat, h_t, cache)
}

///|
/// T-step sequence forward. `input_seq` flat row-major
/// `[seq_len × input_dim]`. Returns `(y_hat_seq, h_final)` where
/// `y_hat_seq` is flat `[seq_len × output_dim]`.
pub fn gru_forecaster_seq_forward(
  model : GRUForecaster,
  input_seq : Array[Float],
  seq_len : Int,
  h_init : Array[Float],
) -> (Array[Float], Array[Float]) {
  let y_hat_seq : Array[Float] = Array::make(
    seq_len * model.output_dim, 0.0F,
  )
  let mut h = h_init
  for t in 0.. Float {
  let n = seq_len * output_dim
  if n <= 0 {
    return 0.0F
  }
  let mut sum_sq = 0.0F
  for i in 0.. Array[Float] {
  let n = seq_len * output_dim
  let grad : Array[Float] = Array::make(n, 0.0F)
  if n <= 0 {
    return grad
  }
  let scale = 2.0F / Float::from_int(n)
  for i in 0.. (GRUForecaster, Float) {
  let in_dim = model.input_dim
  let h_dim = model.hidden_dim
  let out_dim = model.output_dim
  // 1. Forward + per-step caches.
  let y_hat_seq : Array[Float] = Array::make(
    seq_len * out_dim, 0.0F,
  )
  let h_seq : Array[Float] = Array::make(seq_len * h_dim, 0.0F)
  let mut h = h_init
  let gru_cache_seq : Array[GruCellCache] = Array::make(seq_len, {
    x: Array::make(in_dim, 0.0F),
    h_prev: Array::make(h_dim, 0.0F),
    z: Array::make(h_dim, 0.0F),
    r: Array::make(h_dim, 0.0F),
    s: Array::make(h_dim, 0.0F),
    n: Array::make(h_dim, 0.0F),
    h_t: Array::make(h_dim, 0.0F),
  })
  for t in 0..