// lstm_forecaster.mbt — LSTM-based time-series forecaster with full
// BPTT + SGD training (v0.86.0).
//
// Architecture (univariate input, but input_dim is configurable for
// multivariate):
//
//   input_t (length input_dim)
//       ↓ Linear_in (input_dim → hidden_dim)
//   x_proj_t (length hidden_dim)
//       ↓ LstmCell (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} (next-step prediction).
// Training: full BPTT through the LSTM cell (using
// `lstm_cell_backward` from lstm_cell.mbt) plus analytic gradients
// through the input/output Linear projections. Single SGD step on
// all weights per `train_step`.
//
// Reference: Hochreiter & Schmidhuber 1997; identical to the
// input/output pattern of the v0.61.0 LSTM actor but for the
// forecasting use case (scalar regression head instead of tanh
// action squash).

///|
/// LSTM time-series forecaster. Projects input → hidden via Linear_in,
// runs an LstmCell over time, projects hidden → output via Linear_out.
pub struct LSTMForecaster {
  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]
  // LSTM cell
  lstm : LstmCellParam
  // Linear_out: (output_dim × hidden_dim) + bias of length output_dim
  out_w : Array[Array[Float]]
  out_b : Array[Float]
}

///|
/// Build fresh LSTMForecaster. Weights init via xavier_normal scaled
/// by sqrt(2/fan_in) (He-style for ReLU-style activation; here we use
/// tanh internally via the LSTM cell, so this matches the standard).
pub fn LSTMForecaster::new(
  input_dim : Int,
  hidden_dim : Int,
  output_dim : Int,
  seed : UInt64,
) -> LSTMForecaster {
  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 lstm = LstmCellParam::new(hidden_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,
    lstm,
    out_w,
    out_b,
  }
}

///|
/// Per-step forward. Returns `(y_hat, h_t, c_t, cache)`.
/// `y_hat` is the forecast vector (length output_dim) for the NEXT
/// step given the current input. `cache` is the LstmCellCache for
/// BPTT backward.
fn lstm_forecaster_step_with_cache(
  model : LSTMForecaster,
  input_t : Array[Float],
  h_prev : Array[Float],
  c_prev : Array[Float],
) -> (Array[Float], Array[Float], Array[Float], LstmCellCache) {
  // x_proj = in_w · input + in_b
  let x_proj = matvec(model.in_w, model.in_b, input_t)
  // h_t, cache = LSTM_cell(x_proj, h_prev, c_prev)
  let (h_t, cache) = lstm_cell_forward(x_proj, h_prev, c_prev, model.lstm)
  let c_t = cache.c_t
  // y_hat = out_w · h_t + out_b
  let y_hat = matvec(model.out_w, model.out_b, h_t)
  (y_hat, h_t, c_t, cache)
}

///|
/// T-step sequence forward. `input_seq` flat row-major
/// `[seq_len × input_dim]`. Returns `(y_hat_seq, h_final, c_final)`
/// where `y_hat_seq` is flat `[seq_len × output_dim]`.
/// Each step's prediction is the "next step" forecast given the
/// current input — caller typically pairs `y_hat_seq[t]` with
/// `target_seq[t+1]` for training.
pub fn lstm_forecaster_seq_forward(
  model : LSTMForecaster,
  input_seq : Array[Float],
  seq_len : Int,
  h_init : Array[Float],
  c_init : Array[Float],
) -> (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
  let mut c = c_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.. (LSTMForecaster, 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 x_proj_seq : Array[Float] = Array::make(seq_len * h_dim, 0.0F)
  let h_seq : Array[Float] = Array::make(seq_len * h_dim, 0.0F)
  let mut h = h_init
  let mut c = c_init
  let lstm_cache_seq : Array[LstmCellCache] = Array::make(seq_len, {
    x: Array::make(h_dim, 0.0F),
    h_prev: Array::make(h_dim, 0.0F),
    c_prev: Array::make(h_dim, 0.0F),
    f: Array::make(h_dim, 0.0F),
    ig: Array::make(h_dim, 0.0F),
    c_tilde: Array::make(h_dim, 0.0F),
    o: Array::make(h_dim, 0.0F),
    c_t: Array::make(h_dim, 0.0F),
    tanh_c_t: Array::make(h_dim, 0.0F),
    h_t: Array::make(h_dim, 0.0F),
  })
  let c_seq : Array[Float] = Array::make(seq_len * h_dim, 0.0F)
  for t in 0.. d_out_w[i, k] += d_y_hat[t, i] · h_seq[t, k]
  //                    d_out_b[i]    += d_y_hat[t, i]
  let d_out_w : Array[Array[Float]] = Array::make(
    out_dim, Array::make(h_dim, 0.0F),
  )
  let d_out_b : Array[Float] = Array::make(out_dim, 0.0F)
  for t in 0..