// stacked_lstm.mbt — Multi-layer (stacked) LSTM (v0.32.0).
//
// Stacks L LSTM cells where the output hidden state of layer l
// becomes the input of layer l+1 (same timestep).
//
// Conventions:
//   - All layers share the same d_h.
//   - Layer 0 input dim = original d_x.
//   - Layer l > 0 input dim = d_h (same as previous layer output).
//   - The loss is on the top layer's hidden state.
//
// Forward (at each timestep t, layer l):
//   if l == 0: input_l = xs[t]
//   else:      input_l = h_t[l-1]
//   h_t[l], c_t[l] = lstm_cell_forward(input_l, h_{t-1}[l], c_{t-1}[l], params[l])
//
// Backward (BPTT through both layers and time):
//   Walk timesteps backward. At each (t, l), the cell backward
//   receives d_h_t[l] (sum of: d_h from t+1 backward + d_x from
//   layer l+1's backward at time t) and d_c_t[l] (from t+1 backward).
//   The cell returns (d_h_prev, d_x, d_c_prev); d_x feeds back into
//   d_h_t[l-1] (or d_xs[t] if l == 0) at the same timestep, and
//   d_h_prev/d_c_prev are saved as d_h_t/d_c_t for time t-1.

///|
/// Parameter bundle for a stack of L LSTM cells. All layers share
/// d_h; layer 0 has input dim `d_x`, subsequent layers have input
/// dim `d_h`.
pub struct StackedLstmParam {
  d_x : Int
  d_h : Int
  n_layers : Int
  layers : Array[LstmCellParam]
}

///|
/// Build a multi-layer LSTM with `n_layers` cells, each with hidden
/// dim `d_h`. Seeds are layered (seed, seed+1, ...).
pub fn StackedLstmParam::new(
  d_x : Int,
  d_h : Int,
  n_layers : Int,
  seed : UInt64,
) -> StackedLstmParam {
  let layers : Array[LstmCellParam] = []
  for l in 0..