use crate::common::diagnostics::NamErrorCode;
use crate::math::common::AlignedVec;
pub struct LstmLayerDyn {
pub input_size: usize,
pub hidden_size: usize,
pub input_hidden_weights: AlignedVec<f32>,
pub bias: AlignedVec<f32>,
pub state: AlignedVec<f32>,
pub cell_state: AlignedVec<f32>,
pub cell_error: AlignedVec<f32>,
pub gates: AlignedVec<f32>,
}
impl LstmLayerDyn {
pub fn new(input_size: usize, hidden_size: usize) -> Result<Self, NamErrorCode> {
let ih = input_size + hidden_size;
let h4 = 4 * hidden_size;
let weights_len = 4 * ih * hidden_size;
Ok(Self {
input_size,
hidden_size,
input_hidden_weights: AlignedVec::new(weights_len, 0.0f32)?,
bias: AlignedVec::new(h4, 0.0f32)?,
state: AlignedVec::new(ih, 0.0f32)?,
cell_state: AlignedVec::new(hidden_size, 0.0f32)?,
cell_error: AlignedVec::new(hidden_size, 0.0f32)?,
gates: AlignedVec::new(h4, 0.0f32)?,
})
}
#[inline(always)]
pub fn get_hidden_state(&self) -> &[f32] {
&self.state[self.input_size..]
}
pub fn reset_input_slot(&mut self) {
self.state[..self.input_size].fill(0.0);
}
pub fn reset_states(&mut self) {
self.state.fill(0.0);
self.cell_state.fill(0.0);
self.cell_error.fill(0.0);
self.gates.fill(0.0);
}
}