// lstm_cell.mbt — Vanilla LSTM cell forward + backward (v0.29.0).
//
// Single LSTM cell step. Inputs at time t:
// - x_t : input, length d_x
// - h_{t-1} : hidden, length d_h
// - c_{t-1} : cell state, length d_h
//
// Vanilla LSTM (no peepholes):
// f_t = σ(W_f · [h_{t-1}; x_t] + b_f) forget gate
// i_t = σ(W_i · [h_{t-1}; x_t] + b_i) input gate
// C̃_t = tanh(W_C · [h_{t-1}; x_t] + b_C) candidate cell
// C_t = f_t ⊙ C_{t-1} + i_t ⊙ C̃_t cell state
// o_t = σ(W_o · [h_{t-1}; x_t] + b_o) output gate
// h_t = o_t ⊙ tanh(C_t) hidden output
//
// This module provides:
// - `LstmCellParam` — (w_f, w_i, w_c, w_o, b_f, b_i, b_c, b_o)
// - `LstmCellCache` — intermediates for BPTT
// - `LstmCellGrad` — parameter gradients (accumulated over time)
// - `lstm_cell_forward` — single step forward
// - `lstm_cell_backward` — single step backward
// - `lstm_cell_sgd_step` — apply accumulated gradients (SGD)
//
// Note: struct fields MUST be lowercase (MoonBit requirement), so we
// use `w_f` instead of `W_f`. The mathematical notation `W_f` is
// preserved in the doc comments.
///|
/// LSTM cell parameter bundle. Each `w_X` is d_h × (d_h + d_x); each
/// `b_X` has length d_h.
pub struct LstmCellParam {
d_x : Int
d_h : Int
w_f : Array[Array[Float]]
w_i : Array[Array[Float]]
w_c : Array[Array[Float]]
w_o : Array[Array[Float]]
b_f : Array[Float]
b_i : Array[Float]
b_c : Array[Float]
b_o : Array[Float]
}
///|
/// Build a fresh LSTM cell parameter with Xavier-normal init for
/// the four weight matrices and zero biases. seed controls RNG.
pub fn LstmCellParam::new(
d_x : Int,
d_h : Int,
seed : UInt64,
) -> LstmCellParam {
let in_dim = d_h + d_x
let std = sqrtf(2.0F / Float::from_int(d_h + in_dim))
let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
let wf = xavier_normal(d_h, in_dim, std, rng)
let wi = xavier_normal(d_h, in_dim, std, rng)
let wc = xavier_normal(d_h, in_dim, std, rng)
let wo = xavier_normal(d_h, in_dim, std, rng)
let bf : Array[Float] = Array::make(d_h, 0.0F)
let bi : Array[Float] = Array::make(d_h, 0.0F)
let bc : Array[Float] = Array::make(d_h, 0.0F)
let bo : Array[Float] = Array::make(d_h, 0.0F)
{ d_x, d_h, w_f: wf, w_i: wi, w_c: wc, w_o: wo, b_f: bf, b_i: bi, b_c: bc, b_o: bo }
}
///|
/// Cache of intermediate values from a single LSTM forward step.
pub struct LstmCellCache {
x : Array[Float]
h_prev : Array[Float]
c_prev : Array[Float]
f : Array[Float]
ig : Array[Float]
c_tilde : Array[Float]
o : Array[Float]
c_t : Array[Float]
tanh_c_t : Array[Float]
h_t : Array[Float]
}
///|
/// Accumulated parameter gradients from a sequence of backward
/// steps. Initialise all entries to zero and accumulate via
/// `lstm_cell_backward` + `lstm_cell_sgd_step`.
pub struct LstmCellGrad {
d_w_f : Array[Array[Float]]
d_w_i : Array[Array[Float]]
d_w_c : Array[Array[Float]]
d_w_o : Array[Array[Float]]
d_b_f : Array[Float]
d_b_i : Array[Float]
d_b_c : Array[Float]
d_b_o : Array[Float]
}
///|
/// Create a zero-initialised LstmCellGrad matching the param shape.
pub fn LstmCellGrad::zero(param : LstmCellParam) -> LstmCellGrad {
let d_h = param.d_h
let in_dim = param.d_h + param.d_x
let zwf : Array[Array[Float]] = Array::make(d_h, [])
let zwi : Array[Array[Float]] = Array::make(d_h, [])
let zwc : Array[Array[Float]] = Array::make(d_h, [])
let zwo : Array[Array[Float]] = Array::make(d_h, [])
for i in 0.. Array[Float] {
let n = a.length()
let out : Array[Float] = Array::make(n, 0.0F)
for i in 0.. Array[Array[Float]] {
let w : Array[Array[Float]] = Array::make(rows, [])
for i in 0.. Array[Float] {
let n = x.length()
let out : Array[Float] = Array::make(n, 0.0F)
for i in 0.. Array[Float] {
let n = x.length()
let out : Array[Float] = Array::make(n, 0.0F)
for i in 0.. Array[Float] {
let d_out = w.length()
let y : Array[Float] = Array::make(d_out, 0.0F)
for i in 0.. Array[Float] {
let d_out = w.length()
let d_in = x.length()
let d_x : Array[Float] = Array::make(d_in, 0.0F)
for i in 0.. Array[Float] {
let out : Array[Float] = Array::make(a.length() + b.length(), 0.0F)
for i in 0.. (Array[Float], Array[Float]) {
let first : Array[Float] = Array::make(n_first, 0.0F)
let second : Array[Float] = Array::make(v.length() - n_first, 0.0F)
for i in 0.. (Array[Float], LstmCellCache) {
let d_h = param.d_h
let cat = vec_concat(h_prev, x)
let f_pre = matvec(param.w_f, param.b_f, cat)
let f = sigmoid_vec(f_pre)
let i_pre = matvec(param.w_i, param.b_i, cat)
let ig = sigmoid_vec(i_pre)
let c_pre = matvec(param.w_c, param.b_c, cat)
let c_tilde = tanh_vec(c_pre)
let o_pre = matvec(param.w_o, param.b_o, cat)
let o = sigmoid_vec(o_pre)
// Cell state.
let c_t : Array[Float] = Array::make(d_h, 0.0F)
for i in 0.. (Array[Float], Array[Float], Array[Float]) {
let d_h = param.d_h
// d_o = d_h ⊙ tanh(C_t)
let d_o : Array[Float] = Array::make(d_h, 0.0F)
for i in 0.. Unit {
for i in 0..