use std::path::Path;
use crate::{ComputeUnits, DataType, FeatureInfo, Model, MultiArray};
use crate::audio::vad::error::{
ChunkTooLong, ContractMismatch, InferError, ModelError, NonFiniteOutput, OutputShape,
};
pub const CHUNK_SAMPLES: usize = 4096;
pub const CONTEXT_SAMPLES: usize = 64;
pub const MODEL_INPUT_SAMPLES: usize = CONTEXT_SAMPLES + CHUNK_SAMPLES;
pub const STATE_SIZE: usize = 128;
mod names {
pub const AUDIO_INPUT: &str = "audio_input";
pub const HIDDEN_STATE: &str = "hidden_state";
pub const CELL_STATE: &str = "cell_state";
pub const VAD_OUTPUT: &str = "vad_output";
pub const NEW_HIDDEN_STATE: &str = "new_hidden_state";
pub const NEW_CELL_STATE: &str = "new_cell_state";
}
pub const DEFAULT_VAD_COMPUTE: ComputeUnits = ComputeUnits::All;
fn describe(shape: &[usize], dtype: Option<DataType>) -> String {
let dtype = dtype.map_or("none", |d| d.as_str());
format!("{shape:?} {dtype}")
}
#[cfg(feature = "serde")]
fn default_vad_compute() -> ComputeUnits {
DEFAULT_VAD_COMPUTE
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct VadModelOptions {
#[cfg_attr(feature = "serde", serde(default = "default_vad_compute"))]
compute: ComputeUnits,
}
impl Default for VadModelOptions {
fn default() -> Self {
Self::new()
}
}
impl VadModelOptions {
pub const fn new() -> Self {
Self {
compute: DEFAULT_VAD_COMPUTE,
}
}
#[inline(always)]
pub const fn compute(&self) -> ComputeUnits {
self.compute
}
#[must_use]
#[inline(always)]
pub const fn with_compute(mut self, compute: ComputeUnits) -> Self {
self.set_compute(compute);
self
}
#[inline(always)]
pub const fn set_compute(&mut self, compute: ComputeUnits) -> &mut Self {
self.compute = compute;
self
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct VadState {
hidden: [f32; STATE_SIZE],
cell: [f32; STATE_SIZE],
context: [f32; CONTEXT_SAMPLES],
}
impl VadState {
pub const fn initial() -> Self {
Self {
hidden: [0.0; STATE_SIZE],
cell: [0.0; STATE_SIZE],
context: [0.0; CONTEXT_SAMPLES],
}
}
pub const fn from_parts(
hidden: [f32; STATE_SIZE],
cell: [f32; STATE_SIZE],
context: [f32; CONTEXT_SAMPLES],
) -> Self {
Self {
hidden,
cell,
context,
}
}
#[inline(always)]
pub const fn hidden(&self) -> &[f32; STATE_SIZE] {
&self.hidden
}
#[inline(always)]
pub const fn cell(&self) -> &[f32; STATE_SIZE] {
&self.cell
}
#[inline(always)]
pub const fn context(&self) -> &[f32; CONTEXT_SAMPLES] {
&self.context
}
}
impl Default for VadState {
fn default() -> Self {
Self::initial()
}
}
#[derive(Debug)]
pub struct VadModel {
model: Model,
state: VadState,
}
impl VadModel {
pub fn load(path: impl AsRef<Path>) -> Result<Self, ModelError> {
Self::load_with(path, VadModelOptions::new())
}
pub fn load_with(path: impl AsRef<Path>, options: VadModelOptions) -> Result<Self, ModelError> {
let model = Model::load(path, options.compute())?;
let description = model.description();
check_feature(
description.input(names::AUDIO_INPUT),
names::AUDIO_INPUT,
&[1, MODEL_INPUT_SAMPLES],
)?;
check_feature(
description.input(names::HIDDEN_STATE),
names::HIDDEN_STATE,
&[1, STATE_SIZE],
)?;
check_feature(
description.input(names::CELL_STATE),
names::CELL_STATE,
&[1, STATE_SIZE],
)?;
check_feature(
description.output(names::VAD_OUTPUT),
names::VAD_OUTPUT,
&[1, 1, 1],
)?;
check_feature(
description.output(names::NEW_HIDDEN_STATE),
names::NEW_HIDDEN_STATE,
&[1, STATE_SIZE],
)?;
check_feature(
description.output(names::NEW_CELL_STATE),
names::NEW_CELL_STATE,
&[1, STATE_SIZE],
)?;
Ok(Self {
model,
state: VadState::initial(),
})
}
#[inline(always)]
pub const fn state(&self) -> &VadState {
&self.state
}
pub fn reset(&mut self) {
self.state = VadState::initial();
}
pub fn predict_chunk(&mut self, chunk: &[f32]) -> Result<f32, InferError> {
let (probability, next) = self.predict_chunk_with_state(chunk, &self.state)?;
self.state = next;
Ok(probability)
}
pub fn predict_chunk_with_state(
&self,
chunk: &[f32],
state: &VadState,
) -> Result<(f32, VadState), InferError> {
let padded = prepare_chunk(chunk)?;
let window = assemble_window(&state.context, &padded);
check_finite_input(&window)?;
let audio = MultiArray::from_slice(&[1, MODEL_INPUT_SAMPLES], &window)?;
let hidden = MultiArray::from_slice(&[1, STATE_SIZE], &state.hidden)?;
let cell = MultiArray::from_slice(&[1, STATE_SIZE], &state.cell)?;
let mut outputs = self.model.predict_with(&[
(names::AUDIO_INPUT, &audio),
(names::HIDDEN_STATE, &hidden),
(names::CELL_STATE, &cell),
])?;
let probability = take_scalar(&mut outputs, names::VAD_OUTPUT)?;
let next_hidden = take_state(&mut outputs, names::NEW_HIDDEN_STATE)?;
let next_cell = take_state(&mut outputs, names::NEW_CELL_STATE)?;
Ok((
probability,
VadState {
hidden: next_hidden,
cell: next_cell,
context: next_context(&padded),
},
))
}
}
fn check_feature(
feature: Option<&FeatureInfo>,
name: &'static str,
expected_shape: &[usize],
) -> Result<(), ModelError> {
let expected = describe(expected_shape, Some(DataType::F32));
let Some(feature) = feature else {
return Err(ModelError::ContractMismatch(ContractMismatch::new(
name,
expected,
"missing".to_string(),
)));
};
if feature.shape() != expected_shape || feature.data_type() != Some(DataType::F32) {
return Err(ModelError::ContractMismatch(ContractMismatch::new(
name,
expected,
describe(feature.shape(), feature.data_type()),
)));
}
Ok(())
}
fn prepare_chunk(chunk: &[f32]) -> Result<[f32; CHUNK_SAMPLES], InferError> {
if chunk.len() > CHUNK_SAMPLES {
return Err(InferError::ChunkTooLong(ChunkTooLong::new(
chunk.len(),
CHUNK_SAMPLES,
)));
}
let mut padded = [0.0f32; CHUNK_SAMPLES];
padded[..chunk.len()].copy_from_slice(chunk);
let last = chunk.last().copied().unwrap_or(0.0);
for slot in &mut padded[chunk.len()..] {
*slot = last;
}
Ok(padded)
}
fn assemble_window(
context: &[f32; CONTEXT_SAMPLES],
chunk: &[f32; CHUNK_SAMPLES],
) -> [f32; MODEL_INPUT_SAMPLES] {
let mut window = [0.0f32; MODEL_INPUT_SAMPLES];
window[..CONTEXT_SAMPLES].copy_from_slice(context);
window[CONTEXT_SAMPLES..].copy_from_slice(chunk);
window
}
fn next_context(chunk: &[f32; CHUNK_SAMPLES]) -> [f32; CONTEXT_SAMPLES] {
let mut context = [0.0f32; CONTEXT_SAMPLES];
context.copy_from_slice(&chunk[CHUNK_SAMPLES - CONTEXT_SAMPLES..]);
context
}
fn check_finite_input(window: &[f32]) -> Result<(), InferError> {
if let Some(index) = window.iter().position(|v| !v.is_finite()) {
return Err(InferError::NonFiniteInput(index));
}
Ok(())
}
fn take_scalar(outputs: &mut crate::Features, name: &'static str) -> Result<f32, InferError> {
let tensor = outputs
.take(name)
.ok_or_else(|| crate::PredictionError::MissingOutput(name.to_string()))?;
check_output_shape(tensor.shape(), name, &[1, 1, 1])?;
let mut buf = [0.0f32; 1];
tensor.copy_into::<f32>(&mut buf)?;
if !buf[0].is_finite() {
return Err(InferError::NonFiniteOutput(NonFiniteOutput::new(name, 0)));
}
Ok(buf[0])
}
fn take_state(
outputs: &mut crate::Features,
name: &'static str,
) -> Result<[f32; STATE_SIZE], InferError> {
let tensor = outputs
.take(name)
.ok_or_else(|| crate::PredictionError::MissingOutput(name.to_string()))?;
check_output_shape(tensor.shape(), name, &[1, STATE_SIZE])?;
let mut buf = [0.0f32; STATE_SIZE];
tensor.copy_into::<f32>(&mut buf)?;
if let Some(index) = buf.iter().position(|v| !v.is_finite()) {
return Err(InferError::NonFiniteOutput(NonFiniteOutput::new(
name, index,
)));
}
Ok(buf)
}
fn check_output_shape(
shape: &[usize],
feature: &'static str,
expected: &[usize],
) -> Result<(), InferError> {
if shape != expected {
return Err(InferError::OutputShape(OutputShape::new(
feature,
shape.to_vec(),
expected.to_vec(),
)));
}
Ok(())
}
#[cfg(test)]
mod tests;