// 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.. (Array[Array[Array[Float]]], Array[Array[Array[Float]]], StackedLstmCache) {
let n = xs.length()
let l_count = param.n_layers
let hs : Array[Array[Array[Float]]] = []
let cs : Array[Array[Array[Float]]] = []
let caches : Array[Array[LstmCellCache]] = []
let mut cur_h : Array[Array[Float]] = h0s
let mut cur_c : Array[Array[Float]] = c0s
for t in 0.. Float {
let n = hs.length()
if n == 0 {
return 0.0F
}
if target.length() != n {
return -1.0F
}
let mut sum = 0.0F
let mut count = 0
for t in 0.. (Array[Array[Float]], Array[Array[Float]], Array[Array[Float]], Array[LstmCellGrad]) {
let n = hs.length()
let l_count = param.n_layers
let d_h = param.d_h
let d_x = param.d_x
let count = n * d_h
let scale = 2.0F / Float::from_int(count)
// d_xs init.
let d_xs : Array[Array[Float]] = Array::make(n, [])
for t in 0.. Unit {
for l in 0.. Float {
let mut last_loss = 0.0F
for _step in 0.. 0.0F {
for l in 0.. (Array[Array[Float]], Array[Array[Float]]) {
let h0s : Array[Array[Float]] = []
let c0s : Array[Array[Float]] = []
for _ in 0.. (Array[Array[Float]], Array[Array[Float]]) {
let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
let xs : Array[Array[Float]] = Array::make(seq_len, [])
let target : Array[Array[Float]] = Array::make(seq_len, [])
for t in 0.. 0 { t - 1 } else { seq_len - 1 }
let prev = xs[prev_idx]
let overlap = if d_x < d_target { d_x } else { d_target }
for i in 0..