use kopitiam_core::{Error, Result};
use super::Tensor;
#[derive(Debug, Clone)]
pub struct LstmState {
pub cell: Tensor,
pub hidden: Tensor,
}
impl Tensor {
pub fn lstm_cell(
input_gate: &Tensor,
forget_gate: &Tensor,
cell_candidate: &Tensor,
output_gate: &Tensor,
cell_prev: &Tensor,
) -> Result<LstmState> {
let shape = input_gate.shape();
for other in [forget_gate, cell_candidate, output_gate, cell_prev] {
if other.shape() != shape {
return Err(Error::ShapeMismatch {
expected: shape.clone(),
actual: other.shape().clone(),
});
}
}
let i = input_gate.sigmoid()?;
let f = forget_gate.sigmoid()?;
let o = output_gate.sigmoid()?;
let g = cell_candidate.tanh()?;
let cell = f.mul(cell_prev)?.add(&i.mul(&g)?)?;
let hidden = o.mul(&cell.tanh()?)?;
Ok(LstmState { cell, hidden })
}
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_close(a: f32, b: f32) {
assert!((a - b).abs() < 1e-6, "expected {b}, got {a}");
}
#[test]
fn lstm_cell_matches_hand_computation_for_a_single_unit() {
let i_pre = Tensor::from_f32(vec![1.0], [1]).unwrap();
let f_pre = Tensor::from_f32(vec![0.5], [1]).unwrap();
let g_pre = Tensor::from_f32(vec![2.0], [1]).unwrap();
let o_pre = Tensor::from_f32(vec![-0.5], [1]).unwrap();
let c_prev = Tensor::from_f32(vec![0.5], [1]).unwrap();
let out = Tensor::lstm_cell(&i_pre, &f_pre, &g_pre, &o_pre, &c_prev).unwrap();
let i = 1.0f32 / (1.0 + (-1.0f32).exp());
let f = 1.0f32 / (1.0 + (-0.5f32).exp());
let o = 1.0f32 / (1.0 + 0.5f32.exp());
let g = 2.0f32.tanh();
let c = f * 0.5 + i * g;
let h = o * c.tanh();
assert_close(out.cell.to_vec_f32().unwrap()[0], c);
assert_close(out.hidden.to_vec_f32().unwrap()[0], h);
}
#[test]
fn lstm_cell_with_zero_input_and_zero_state_stays_at_zero_cell() {
let z = Tensor::zeros([4]);
let out = Tensor::lstm_cell(&z, &z, &z, &z, &z).unwrap();
assert_eq!(out.cell.to_vec_f32().unwrap(), vec![0.0; 4]);
assert_eq!(out.hidden.to_vec_f32().unwrap(), vec![0.0; 4]);
}
#[test]
fn forget_open_and_input_shut_preserves_the_cell_state_exactly() {
let big = 30.0f32; let f_pre = Tensor::from_f32(vec![big, big, big], [3]).unwrap();
let i_pre = Tensor::from_f32(vec![-big, -big, -big], [3]).unwrap();
let g_pre = Tensor::from_f32(vec![5.0, -5.0, 1.0], [3]).unwrap(); let o_pre = Tensor::from_f32(vec![0.0, 0.0, 0.0], [3]).unwrap();
let c_prev = Tensor::from_f32(vec![0.3, -0.7, 1.2], [3]).unwrap();
let out = Tensor::lstm_cell(&i_pre, &f_pre, &g_pre, &o_pre, &c_prev).unwrap();
let cell = out.cell.to_vec_f32().unwrap();
for (got, want) in cell.iter().zip([0.3, -0.7, 1.2]) {
assert!((got - want).abs() < 1e-6, "cell state not preserved: got {got}, want {want}");
}
}
#[test]
fn lstm_cell_preserves_a_multi_dimensional_shape() {
let a = Tensor::from_f32(vec![0.1, 0.2, 0.3, 0.4], [2, 2]).unwrap();
let out = Tensor::lstm_cell(&a, &a, &a, &a, &a).unwrap();
assert_eq!(out.cell.shape().dims(), &[2, 2]);
assert_eq!(out.hidden.shape().dims(), &[2, 2]);
}
#[test]
fn lstm_cell_rejects_a_shape_disagreement() {
let a = Tensor::from_f32(vec![0.0; 3], [3]).unwrap();
let b = Tensor::from_f32(vec![0.0; 4], [4]).unwrap();
assert!(matches!(
Tensor::lstm_cell(&a, &a, &a, &a, &b),
Err(Error::ShapeMismatch { .. })
));
}
#[test]
fn lstm_cell_rejects_non_f32_input() {
let f32t = Tensor::from_f32(vec![0.0; 3], [3]).unwrap();
let i32t = Tensor::from_i32(vec![0; 3], [3]).unwrap();
assert!(matches!(
Tensor::lstm_cell(&i32t, &f32t, &f32t, &f32t, &f32t),
Err(Error::DTypeMismatch { .. })
));
}
}