// gru_demo.mbt — GRU sequence processing + training demo (v0.30.1).
//
// Sequence-level GRU operations:
// - `gru_sequence_forward` — run a GRU cell over an input sequence,
// returning per-step hidden states and per-step caches for BPTT.
// - `gru_sequence_loss` — mean-squared-error loss between the
// sequence of hidden states and a target sequence.
// - `gru_sequence_backward` — BPTT: walk caches in reverse, sum
// parameter gradients, accumulate per-input gradients.
// - `gru_train_n_steps` — K-step SGD on a single sequence with
// optional gradient clipping (mirrors lstm_train_n_steps).
// - `gru_identity_dataset` — synthetic identity-shift task with
// separate input dim and target dim.
///|
/// Run a GRU cell over an input sequence. Returns
/// (hs, caches) where:
/// - `hs[t]` is the hidden state at time t (length d_h)
/// - `caches` has length `seq_len` for BPTT.
pub fn gru_sequence_forward(
xs : Array[Array[Float]],
h0 : Array[Float],
param : GruCellParam,
) -> (Array[Array[Float]], Array[GruCellCache]) {
let n = xs.length()
let hs : Array[Array[Float]] = []
let caches : Array[GruCellCache] = []
let mut cur_h : Array[Float] = h0
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], GruCellGrad) {
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 {
for i in 0.. Unit {
let d_h = grad.d_w_z.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 gru_train_n_steps(
xs : Array[Array[Float]],
target : Array[Array[Float]],
h0 : Array[Float],
param : GruCellParam,
n_steps : Int,
lr : Float,
clip : Float,
) -> Float {
let mut last_loss = 0.0F
for _step in 0.. 0.0F {
gru_clip_grad(grad, clip)
}
gru_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). xs has dim `d_x`; target has
/// dim `d_target`.
pub fn gru_identity_dataset(
seq_len : Int,
d_x : Int,
d_target : 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 { 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..