// 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.. Float {
let n = predicted.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[Float], Array[Float], LstmCellGrad) {
let n = predicted.length()
let d_h = param.d_h
let d_x = param.d_x
// d_loss / d_h_t = 2 (h_t - target_t) / count
let count = n * d_h
let scale = 2.0F / Float::from_int(count)
let d_xs : Array[Array[Float]] = Array::make(n, [])
for t in 0.. 0 we have to add
// d_h_prev to the previous step's d_h (which is d_xs[t-1]).
let d_h_prev_global : Array[Float] = Array::make(d_h, 0.0F)
for t_step in 0.. 0 propagation) to d_xs[t] before
// running the cell backward.
let d_xs_t : Array[Float] = d_xs[t]
if t > 0 {
for i in 0.. Unit {
let d_h = grad.d_w_f.length()
for i in 0.. Float {
if x > clip {
return clip
}
if x < -clip {
return -clip
}
x
}
///|
/// Run K-step SGD training on a single input/target sequence. Returns
/// the final loss. Optionally clips gradients to `[-clip, clip]`.
pub fn lstm_train_n_steps(
xs : Array[Array[Float]],
target : Array[Array[Float]],
h0 : Array[Float],
c0 : Array[Float],
param : LstmCellParam,
n_steps : Int,
lr : Float,
clip : Float,
) -> Float {
let mut last_loss = 0.0F
for _step in 0.. 0.0F {
lstm_clip_grad(grad, clip)
}
lstm_cell_sgd_step(param, grad, lr)
last_loss = loss
}
last_loss
}
///|
/// Generate a simple synthetic sequence-prediction task: the target
/// at time t+1 is a delayed version of input at time t (i.e., learn
/// the identity over a 1-step shift). Returns (xs, target) both
/// length `seq_len`, each entry length `d_x`.
pub fn lstm_identity_dataset(
seq_len : Int,
d_x : Int,
seed : UInt64,
) -> (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 { xs[t - 1] } else { x }
g[i] = prev[i]
}
xs[t] = x
target[t] = g
}
(xs, target)
}