use std::path::Path;
use ort::session::Session;
use ort::value::{Tensor, TensorRef};
use crate::error::{Result, VadError};
pub const FRAME_LEN: usize = 512;
const STATE_SHAPE: [usize; 3] = [2, 1, 128];
const STATE_LEN: usize = STATE_SHAPE[0] * STATE_SHAPE[1] * STATE_SHAPE[2];
pub struct SileroVad {
session: Session,
state: Vec<f32>,
sr: Tensor<i64>,
}
impl SileroVad {
pub fn new(model_path: &Path) -> Result<Self> {
let mut builder = Session::builder().map_err(|e| VadError::ModelLoad(e.to_string()))?;
let session = builder
.commit_from_file(model_path)
.map_err(|e| VadError::ModelLoad(format!("{}: {e}", model_path.display())))?;
let sr = Tensor::from_array((Vec::<i64>::new(), vec![16_000_i64]))
.map_err(|e| VadError::ModelLoad(format!("failed to build sr tensor: {e}")))?;
Ok(Self {
session,
state: vec![0.0; STATE_LEN],
sr,
})
}
pub fn process(&mut self, frame: &[f32]) -> Result<f32> {
if frame.len() != FRAME_LEN {
return Err(VadError::BadFrameLen {
expected: FRAME_LEN,
actual: frame.len(),
}
.into());
}
let input = TensorRef::from_array_view(([1_usize, FRAME_LEN], frame))
.map_err(|e| VadError::Inference(e.to_string()))?;
let state = TensorRef::from_array_view((STATE_SHAPE, &self.state[..]))
.map_err(|e| VadError::Inference(e.to_string()))?;
let outputs = self
.session
.run(ort::inputs!["input" => input, "state" => state, "sr" => &self.sr])
.map_err(|e| VadError::Inference(e.to_string()))?;
let prob = outputs["output"]
.try_extract_tensor::<f32>()
.map_err(|e| VadError::Inference(e.to_string()))?
.1[0];
let new_state = outputs["stateN"]
.try_extract_tensor::<f32>()
.map_err(|e| VadError::Inference(e.to_string()))?
.1;
self.state.copy_from_slice(new_state);
Ok(prob)
}
pub fn reset_state(&mut self) {
self.state.fill(0.0);
}
}