// gru_demo.mbt — GRU sequence processing + training demo (v0.30.1).
//
// Sequence-level GRU operations:
//   - `gru_sequence_forward` — run a GRU cell over an input sequence,
//     returning per-step hidden states and per-step caches for BPTT.
//   - `gru_sequence_loss` — mean-squared-error loss between the
//     sequence of hidden states and a target sequence.
//   - `gru_sequence_backward` — BPTT: walk caches in reverse, sum
//     parameter gradients, accumulate per-input gradients.
//   - `gru_train_n_steps` — K-step SGD on a single sequence with
//     optional gradient clipping (mirrors lstm_train_n_steps).
//   - `gru_identity_dataset` — synthetic identity-shift task with
//     separate input dim and target dim.

///|
/// Run a GRU cell over an input sequence. Returns
/// (hs, caches) where:
///   - `hs[t]` is the hidden state at time t (length d_h)
///   - `caches` has length `seq_len` for BPTT.
pub fn gru_sequence_forward(
  xs : Array[Array[Float]],
  h0 : Array[Float],
  param : GruCellParam,
) -> (Array[Array[Float]], Array[GruCellCache]) {
  let n = xs.length()
  let hs : Array[Array[Float]] = []
  let caches : Array[GruCellCache] = []
  let mut cur_h : Array[Float] = h0
  for t in 0..