// bilstm.mbt — Bidirectional LSTM (v0.32.1).
//
// Runs two LSTMs over the input:
// - Forward: standard left-to-right
// - Backward: right-to-left
//
// Output hidden state at time t is the concat
// h_combined[t] = [h_f[t]; h_b[t]]
// of length 2 * d_h. The loss is on h_combined.
//
// Backward:
// d_h_combined[t] -> d_h_f[t] (first d_h) + d_h_b[t] (last d_h)
// - Forward LSTM BPTT walks t = T-1, ..., 0.
// - Backward LSTM BPTT walks t = 0, ..., T-1 (i.e., reverses time).
// Both contribute to d_xs[t] which is summed.
///|
/// Parameter bundle for a 1-layer BiLSTM. Each direction is its
/// own LstmCellParam.
pub struct BiLstmParam {
d_x : Int
d_h : Int
fwd : LstmCellParam
bwd : LstmCellParam
}
///|
/// Build a 1-layer BiLSTM with both directions having hidden dim
/// `d_h`. Forward uses seed `seed`, backward uses `seed + 1`.
pub fn BiLstmParam::new(d_x : Int, d_h : Int, seed : UInt64) -> BiLstmParam {
{
d_x,
d_h,
fwd: LstmCellParam::new(d_x, d_h, seed),
bwd: LstmCellParam::new(d_x, d_h, seed + 1UL),
}
}
///|
/// Forward cache storing per-direction per-step cell caches plus
/// the forward outputs (used by BPTT and for loss computation).
pub struct BiLstmCache {
caches_f : Array[LstmCellCache]
caches_b : Array[LstmCellCache]
hs_f : Array[Array[Float]]
hs_b : Array[Array[Float]]
cs_f : Array[Array[Float]]
cs_b : Array[Array[Float]]
}
///|
/// Run the BiLSTM forward over an input sequence. Returns
/// `(hs_f, hs_b, cs_f, cs_b, cache)`. `hs_f[t]` and `hs_b[t]` are
/// the per-direction hidden states (each length d_h).
pub fn bilstm_forward(
xs : Array[Array[Float]],
h0_f : Array[Float],
c0_f : Array[Float],
h0_b : Array[Float],
c0_b : Array[Float],
param : BiLstmParam,
) -> (Array[Array[Float]], Array[Array[Float]], Array[Array[Float]], Array[Array[Float]], BiLstmCache) {
let n = xs.length()
let hs_f : Array[Array[Float]] = []
let cs_f : Array[Array[Float]] = []
let caches_f : Array[LstmCellCache] = []
let hs_b : Array[Array[Float]] = []
let cs_b : Array[Array[Float]] = []
let caches_b : Array[LstmCellCache] = []
// Forward direction: left-to-right.
let mut cur_h = h0_f
let mut cur_c = c0_f
for t in 0..= 0 {
let (h_t, cache_t) = lstm_cell_forward(xs[t], cur_h, cur_c, param.bwd)
hs_b.push(h_t)
cs_b.push(cache_t.c_t)
caches_b.push(cache_t)
cur_h = h_t
cur_c = cache_t.c_t
t = t - 1
}
let cache : BiLstmCache = {
caches_f, caches_b, hs_f, hs_b, cs_f, cs_b,
}
(hs_f, hs_b, cs_f, cs_b, cache)
}
///|
/// Build the concat hidden state `[hs_f[t]; hs_b[t]]` at time t.
pub fn bilstm_concat(hs_f : Array[Array[Float]], hs_b : Array[Array[Float]], t : Int) -> Array[Float] {
let d_h = hs_f[t].length()
let out : Array[Float] = Array::make(2 * d_h, 0.0F)
for i in 0.. Float {
let n = hs_f.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], Array[Float], Array[Float], LstmCellGrad, LstmCellGrad) {
let n = hs_f.length()
let d_h = param.d_h
let d_x = param.d_x
let two_dh = 2 * d_h
let count = n * two_dh
let scale = 2.0F / Float::from_int(count)
let d_xs : Array[Array[Float]] = Array::make(n, [])
for t in 0.. Unit {
lstm_cell_sgd_step(param.fwd, grad_f, lr)
lstm_cell_sgd_step(param.bwd, grad_b, lr)
}
///|
/// Run K-step SGD training on a single input/target sequence.
pub fn bilstm_train_n_steps(
xs : Array[Array[Float]],
target : Array[Array[Float]],
h0_f : Array[Float],
c0_f : Array[Float],
h0_b : Array[Float],
c0_b : Array[Float],
param : BiLstmParam,
n_steps : Int,
lr : Float,
clip : Float,
) -> Float {
let mut last_loss = 0.0F
for _step in 0.. 0.0F {
lstm_clip_grad(grad_f, clip)
lstm_clip_grad(grad_b, clip)
}
bilstm_sgd_step(param, grad_f, grad_b, lr)
last_loss = loss
}
last_loss
}
///|
/// Synthetic identity-shift task for BiLSTM. xs has dim d_x;
/// target has dim 2 * d_target (matching concat hidden dim).
pub fn bilstm_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 next_idx = if t < seq_len - 1 { t + 1 } else { 0 }
let prev = xs[prev_idx]
let next = xs[next_idx]
let overlap = if d_x < d_target { d_x } else { d_target }
for i in 0..