// bilstm.mbt — Bidirectional LSTM (v0.32.1).
//
// Runs two LSTMs over the input:
//   - Forward:  standard left-to-right
//   - Backward: right-to-left
//
// Output hidden state at time t is the concat
//   h_combined[t] = [h_f[t]; h_b[t]]
// of length 2 * d_h. The loss is on h_combined.
//
// Backward:
//   d_h_combined[t] -> d_h_f[t] (first d_h) + d_h_b[t] (last d_h)
//   - Forward LSTM BPTT walks t = T-1, ..., 0.
//   - Backward LSTM BPTT walks t = 0, ..., T-1 (i.e., reverses time).
//   Both contribute to d_xs[t] which is summed.

///|
/// Parameter bundle for a 1-layer BiLSTM. Each direction is its
/// own LstmCellParam.
pub struct BiLstmParam {
  d_x : Int
  d_h : Int
  fwd : LstmCellParam
  bwd : LstmCellParam
}

///|
/// Build a 1-layer BiLSTM with both directions having hidden dim
/// `d_h`. Forward uses seed `seed`, backward uses `seed + 1`.
pub fn BiLstmParam::new(d_x : Int, d_h : Int, seed : UInt64) -> BiLstmParam {
  {
    d_x,
    d_h,
    fwd: LstmCellParam::new(d_x, d_h, seed),
    bwd: LstmCellParam::new(d_x, d_h, seed + 1UL),
  }
}

///|
/// Forward cache storing per-direction per-step cell caches plus
/// the forward outputs (used by BPTT and for loss computation).
pub struct BiLstmCache {
  caches_f : Array[LstmCellCache]
  caches_b : Array[LstmCellCache]
  hs_f : Array[Array[Float]]
  hs_b : Array[Array[Float]]
  cs_f : Array[Array[Float]]
  cs_b : Array[Array[Float]]
}

///|
/// Run the BiLSTM forward over an input sequence. Returns
/// `(hs_f, hs_b, cs_f, cs_b, cache)`. `hs_f[t]` and `hs_b[t]` are
/// the per-direction hidden states (each length d_h).
pub fn bilstm_forward(
  xs : Array[Array[Float]],
  h0_f : Array[Float],
  c0_f : Array[Float],
  h0_b : Array[Float],
  c0_b : Array[Float],
  param : BiLstmParam,
) -> (Array[Array[Float]], Array[Array[Float]], Array[Array[Float]], Array[Array[Float]], BiLstmCache) {
  let n = xs.length()
  let hs_f : Array[Array[Float]] = []
  let cs_f : Array[Array[Float]] = []
  let caches_f : Array[LstmCellCache] = []
  let hs_b : Array[Array[Float]] = []
  let cs_b : Array[Array[Float]] = []
  let caches_b : Array[LstmCellCache] = []
  // Forward direction: left-to-right.
  let mut cur_h = h0_f
  let mut cur_c = c0_f
  for t in 0..= 0 {
    let (h_t, cache_t) = lstm_cell_forward(xs[t], cur_h, cur_c, param.bwd)
    hs_b.push(h_t)
    cs_b.push(cache_t.c_t)
    caches_b.push(cache_t)
    cur_h = h_t
    cur_c = cache_t.c_t
    t = t - 1
  }
  let cache : BiLstmCache = {
    caches_f, caches_b, hs_f, hs_b, cs_f, cs_b,
  }
  (hs_f, hs_b, cs_f, cs_b, cache)
}

///|
/// Build the concat hidden state `[hs_f[t]; hs_b[t]]` at time t.
pub fn bilstm_concat(hs_f : Array[Array[Float]], hs_b : Array[Array[Float]], t : Int) -> Array[Float] {
  let d_h = hs_f[t].length()
  let out : Array[Float] = Array::make(2 * d_h, 0.0F)
  for i in 0..