// gru_cell.mbt — Vanilla GRU cell forward + backward + SGD (v0.30.0).
//
// Single GRU cell step. Inputs at time t:
// - x_t : input, length d_x
// - h_{t-1} : hidden, length d_h
//
// Vanilla GRU (no reset-after-Fourier tricks, no layer-norm):
// z_t = σ(W_z · [h_{t-1}; x_t] + b_z) update gate
// r_t = σ(W_r · [h_{t-1}; x_t] + b_r) reset gate
// s_t = r_t ⊙ h_{t-1} reset-applied hidden
// n_t = tanh(W_n · [s_t; x_t] + b_n) candidate
// h_t = (1 - z_t) ⊙ h_{t-1} + z_t ⊙ n_t hidden output
//
// This module provides:
// - `GruCellParam` — (w_z, w_r, w_n, b_z, b_r, b_n)
// - `GruCellCache` — intermediates for BPTT
// - `GruCellGrad` — parameter gradients (accumulated over time)
// - `gru_cell_forward` — single step forward
// - `gru_cell_backward` — single step backward
// - `gru_cell_sgd_step` — apply accumulated gradients (SGD)
//
// Like LSTM, struct fields MUST be lowercase (MoonBit requirement),
// so we use `w_z` instead of `W_z`. The mathematical notation `W_z`
// is preserved in the doc comments.
///|
/// GRU cell parameter bundle. Each `w_X` is d_h × (d_h + d_x); each
/// `b_X` has length d_h. W_n also has shape d_h × (d_h + d_x), but
/// its first d_h input slots see `r ⊙ h_prev` (not `h_prev`).
pub struct GruCellParam {
d_x : Int
d_h : Int
w_z : Array[Array[Float]]
w_r : Array[Array[Float]]
w_n : Array[Array[Float]]
b_z : Array[Float]
b_r : Array[Float]
b_n : Array[Float]
}
///|
/// Build a fresh GRU cell parameter with Xavier-normal init for the
/// three weight matrices and zero biases. seed controls RNG.
pub fn GruCellParam::new(
d_x : Int,
d_h : Int,
seed : UInt64,
) -> GruCellParam {
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 wz = xavier_normal(d_h, in_dim, std, rng)
let wr = xavier_normal(d_h, in_dim, std, rng)
let wn = xavier_normal(d_h, in_dim, std, rng)
let bz : Array[Float] = Array::make(d_h, 0.0F)
let br : Array[Float] = Array::make(d_h, 0.0F)
let bn : Array[Float] = Array::make(d_h, 0.0F)
{ d_x, d_h, w_z: wz, w_r: wr, w_n: wn, b_z: bz, b_r: br, b_n: bn }
}
///|
/// Cache of intermediate values from a single GRU forward step.
pub struct GruCellCache {
x : Array[Float]
h_prev : Array[Float]
z : Array[Float]
r : Array[Float]
s : Array[Float]
n : Array[Float]
h_t : Array[Float]
}
///|
/// Accumulated parameter gradients from a sequence of backward
/// steps. Initialise all entries to zero and accumulate via
/// `gru_cell_backward` + `gru_cell_sgd_step`.
pub struct GruCellGrad {
d_w_z : Array[Array[Float]]
d_w_r : Array[Array[Float]]
d_w_n : Array[Array[Float]]
d_b_z : Array[Float]
d_b_r : Array[Float]
d_b_n : Array[Float]
}
///|
/// Create a zero-initialised GruCellGrad matching the param shape.
pub fn GruCellGrad::zero(param : GruCellParam) -> GruCellGrad {
let d_h = param.d_h
let in_dim = param.d_h + param.d_x
let zwz : Array[Array[Float]] = Array::make(d_h, [])
let zwr : Array[Array[Float]] = Array::make(d_h, [])
let zwn : Array[Array[Float]] = Array::make(d_h, [])
for i in 0.. (Array[Float], GruCellCache) {
let d_h = param.d_h
let cat = vec_concat(h_prev, x)
let z_pre = matvec(param.w_z, param.b_z, cat)
let z = sigmoid_vec(z_pre)
let r_pre = matvec(param.w_r, param.b_r, cat)
let r = sigmoid_vec(r_pre)
// s = r ⊙ h_prev
let s : Array[Float] = Array::make(d_h, 0.0F)
for i in 0.. (Array[Float], Array[Float]) {
let d_h = param.d_h
let d_x = param.d_x
let in_dim = d_h + d_x
// 1. From h_t = (1-z) ⊙ h_prev + z ⊙ n:
// d_z = d_h ⊙ (n - h_prev)
// d_n = d_h ⊙ z
// d_h_prev_main = d_h ⊙ (1 - z)
let d_z : Array[Float] = Array::make(d_h, 0.0F)
let d_n : Array[Float] = Array::make(d_h, 0.0F)
let d_h_prev_main : Array[Float] = Array::make(d_h, 0.0F)
for i in 0.. Unit {
for i in 0..