use super::layer_dyn::LstmLayerDyn;
use crate::common::diagnostics::NamErrorCode;
use crate::math::common::AlignedVec;
pub struct LstmModelDyn {
pub layers: Vec<LstmLayerDyn>,
pub head_weights: AlignedVec<f32>,
pub head_weights_f32: AlignedVec<f32>,
pub head_bias: f32,
pub prewarm_on_reset: bool,
pub expected_sample_rate: f64,
}
impl LstmModelDyn {
pub fn new(num_layers: usize, hidden_size: usize) -> Result<Self, NamErrorCode> {
let mut layers = Vec::with_capacity(num_layers);
for i in 0..num_layers {
let input_size = if i == 0 { 1 } else { hidden_size };
layers.push(LstmLayerDyn::new(input_size, hidden_size)?);
}
Ok(Self {
layers,
head_weights: AlignedVec::new(hidden_size, 0.0f32)?,
head_weights_f32: AlignedVec::new(hidden_size, 0.0f32)?,
head_bias: 0.0,
prewarm_on_reset: true,
expected_sample_rate: 48000.0,
})
}
#[target_feature(enable = "avx2,fma,f16c")]
unsafe fn process_avx2(&mut self, input: &[f32], output: &mut [f32]) {
if self.layers.is_empty() {
return;
}
unsafe {
let n_layers = self.layers.len();
debug_assert!(n_layers > 0, "LstmModelDyn requires at least one layer");
let layers_ptr = self.layers.as_mut_ptr();
for (s, &val) in input.iter().enumerate() {
(*layers_ptr).process_sample_avx2(&[val]);
for i in 1..n_layers {
let prev = &*layers_ptr.add(i - 1);
let hidden = &prev.state[prev.input_size..];
(*layers_ptr.add(i)).process_sample_avx2(hidden);
}
let last = &*layers_ptr.add(n_layers - 1);
let h = last.get_hidden_state();
let dot = crate::math::common::scalar_ref::dot_product_f32_native_kahan(
h,
&self.head_weights_f32,
);
output[s] = dot + self.head_bias;
}
}
}
#[target_feature(enable = "avx512f,avx512vl")]
unsafe fn process_avx512(&mut self, input: &[f32], output: &mut [f32]) {
if self.layers.is_empty() {
return;
}
unsafe {
let n_layers = self.layers.len();
debug_assert!(n_layers > 0, "LstmModelDyn requires at least one layer");
let layers_ptr = self.layers.as_mut_ptr();
for (s, &val) in input.iter().enumerate() {
(*layers_ptr).process_sample_avx512(&[val]);
for i in 1..n_layers {
let prev = &*layers_ptr.add(i - 1);
let hidden = &prev.state[prev.input_size..];
(*layers_ptr.add(i)).process_sample_avx512(hidden);
}
let last = &*layers_ptr.add(n_layers - 1);
let h = last.get_hidden_state();
let dot = crate::math::common::scalar_ref::dot_product_f32_native_kahan(
h,
&self.head_weights_f32,
);
output[s] = dot + self.head_bias;
}
}
}
pub fn process(&mut self, input: &[f32], output: &mut [f32]) {
unsafe {
crate::math::common::dispatch_simd!(
@self,
process_avx512,
process_avx512,
process_avx2,
input,
output
);
}
}
pub fn process_scalar(&mut self, input: &[f32], output: &mut [f32]) {
if self.layers.is_empty() {
return;
}
let n_layers = self.layers.len();
debug_assert!(n_layers > 0, "LstmModelDyn requires at least one layer");
let layers_ptr = self.layers.as_mut_ptr();
for s in 0..input.len() {
unsafe {
(*layers_ptr).process_sample_scalar(&[input[s]]);
for i in 1..n_layers {
let prev = &*layers_ptr.add(i - 1);
let hidden = &prev.state[prev.input_size..];
let hidden_copy: Vec<f32> = hidden.to_vec();
(*layers_ptr.add(i)).process_sample_scalar(&hidden_copy);
}
let last = &*layers_ptr.add(n_layers - 1);
let hidden_last = last.get_hidden_state();
let dot = crate::math::common::scalar_ref::dot_product_f32_native_kahan(
hidden_last,
&self.head_weights_f32,
);
output[s] = dot + self.head_bias;
}
}
}
pub fn reset_states(&mut self) {
for layer in &mut self.layers {
layer.reset_states();
}
}
pub fn reset_input_slots(&mut self) {
for layer in &mut self.layers {
layer.reset_input_slot();
}
}
}