// lstm_demo.mbt — LSTM sequence processing + training demo (v0.29.1).
//
// Sequence-level LSTM operations:
//   - `lstm_sequence_forward` — run an LSTM cell over an input sequence,
///     returning per-step hidden states and per-step caches for BPTT.
//   - `lstm_sequence_loss` — mean-squared-error loss between the
///     sequence of hidden states and a target sequence.
//   - `lstm_sequence_backward` — BPTT: walk caches in reverse, sum
///     parameter gradients, accumulate per-input gradients.
//   - `lstm_train_n_steps` — K-step SGD on a single sequence with
///     optional gradient clipping (mirrors scnn_train_n_steps).

///|
/// Run an LSTM cell over an input sequence. Returns
/// (hs, cs_final, caches) where:
///   - `hs[t]` is the hidden state at time t (length d_h)
///   - `cs_final` is the final cell state (length d_h)
///   - `caches` has length `seq_len` for BPTT.
pub fn lstm_sequence_forward(
  xs : Array[Array[Float]],
  h0 : Array[Float],
  c0 : Array[Float],
  param : LstmCellParam,
) -> (Array[Array[Float]], Array[Float], Array[LstmCellCache]) {
  let n = xs.length()
  let d_h = param.d_h
  let hs : Array[Array[Float]] = []
  let caches : Array[LstmCellCache] = []
  let mut cur_h : Array[Float] = h0
  let mut cur_c : Array[Float] = c0
  for t in 0..