use std::collections::HashMap;
use smol_str::format_smolstr;
use crate::{
array::Array,
audio::vad::{
models::silero_vad::config::{BranchConfig, ModelConfig},
output::SpeechSegment,
},
dtype::Dtype,
error::{Error, OutOfRangePayload, RankMismatchPayload, Result},
ops,
};
const EVAL_EVERY: usize = 16;
pub struct SileroVadBranch {
config: BranchConfig,
stft_conv_weight: Array,
conv1: ConvBlock,
conv2: ConvBlock,
conv3: ConvBlock,
conv4: ConvBlock,
lstm: Lstm,
final_conv_weight: Array,
final_conv_bias: Array,
}
struct ConvBlock {
weight: Array,
bias: Array,
stride: i32,
padding: i32,
}
impl ConvBlock {
fn forward(&self, x: &Array) -> Result<Array> {
let h = ops::conv::conv1d(x, &self.weight, self.stride, self.padding, 1, 1)?;
let h = h.add(&self.bias)?;
relu(&h)
}
}
fn scalar_f32(value: f32) -> Result<Array> {
Array::from_slice::<f32>(&[value], &[0i32; 0])
}
fn idx0(value: i32) -> Result<Array> {
Array::from_slice::<i32>(&[value], &[0i32; 0])
}
fn relu(x: &Array) -> Result<Array> {
let zero = scalar_f32(0.0)?.astype(x.dtype()?)?;
ops::arithmetic::maximum(x, &zero)
}
fn reflect_pad_right(x: &Array, pad: i32) -> Result<Array> {
if pad <= 0 {
return x.try_clone();
}
let shape = x.shape();
let last = *shape.last().ok_or_else(|| {
Error::RankMismatch(RankMismatchPayload::new(
"reflect_pad_right: input",
shape.len() as u32,
shape.clone(),
))
})?;
let len = i32::try_from(last).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"reflect_pad_right: input length",
"must fit in i32",
format_smolstr!("{last}"),
))
})?;
if len <= pad {
return Err(Error::OutOfRange(OutOfRangePayload::new(
"reflect_pad_right: pad vs samples",
"reflect padding requires more samples than the pad width",
format_smolstr!("pad={pad}, samples={len}"),
)));
}
let indices = Array::arange::<i32>(len - 2, len - pad - 2, -1)?;
let reflected = ops::indexing::take_axis(x, &indices, -1)?;
ops::shape::concatenate(&[x, &reflected], -1)
}
impl SileroVadBranch {
#[inline(always)]
pub const fn config(&self) -> &BranchConfig {
&self.config
}
pub fn forward(&self, x: &Array, state: Option<&Array>) -> Result<(Array, Array)> {
let promoted;
let x = if x.ndim() == 1 {
promoted = ops::shape::expand_dims_axes(x, &[0])?;
&promoted
} else {
x
};
let (hidden, cell) = split_state(state)?;
let x = reflect_pad_right(x, self.config.pad())?;
let x = ops::shape::expand_dims_axes(&x, &[-1])?;
let x = ops::conv::conv1d(
&x,
&self.stft_conv_weight,
self.config.hop_length(),
0,
1,
1,
)?;
let cutoff = self.config.cutoff();
let real = slice_last_axis(&x, 0, cutoff)?;
let imag = slice_last_axis(&x, cutoff, 2 * cutoff)?;
let mag = real.multiply(&real)?.add(&imag.multiply(&imag)?)?;
let x = ops::arithmetic::sqrt(&mag)?;
let x = self.conv1.forward(&x)?;
let x = self.conv2.forward(&x)?;
let x = self.conv3.forward(&x)?;
let x = self.conv4.forward(&x)?;
let (hidden_seq, cell_seq) = self.lstm.forward(&x, hidden.as_ref(), cell.as_ref())?;
let last_hidden = last_timestep(&hidden_seq)?;
let last_cell = last_timestep(&cell_seq)?;
let new_state = ops::shape::stack_axis(&[&last_hidden, &last_cell], 0)?;
let x = relu(&hidden_seq)?;
let x = ops::conv::conv1d(&x, &self.final_conv_weight, 1, 0, 1, 1)?;
let x = x.add(&self.final_conv_bias)?;
let x = ops::arithmetic::sigmoid(&x)?;
let x = ops::shape::squeeze_axes(&x, &[-1])?;
let x = ops::reduction::mean_axes(&x, &[1], true)?;
Ok((x, new_state))
}
}
fn split_state(state: Option<&Array>) -> Result<(Option<Array>, Option<Array>)> {
let Some(state) = state else {
return Ok((None, None));
};
let shape = state.shape();
if shape.len() != 3 || shape[0] != 2 {
let rank = shape.len() as u32;
return Err(Error::RankMismatch(RankMismatchPayload::new(
"silero_vad state: expected (2, batch, 128)",
rank,
shape,
)));
}
let hidden = state.take_axis(&idx0(0)?, 0)?;
let cell = state.take_axis(&idx0(1)?, 0)?;
Ok((Some(hidden), Some(cell)))
}
fn last_timestep(seq: &Array) -> Result<Array> {
let l = i32::try_from(seq.shape()[1]).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"silero_vad lstm seq length",
"must fit in i32",
format_smolstr!("{}", seq.shape()[1]),
))
})?;
let idx = idx0(l - 1)?;
seq.take_axis(&idx, 1)
}
fn slice_last_axis(x: &Array, lo: i32, hi: i32) -> Result<Array> {
let shape = x.shape();
let n = shape.len();
let mut start = vec![0_i32; n];
let mut stop: Vec<i32> = shape
.iter()
.map(|&d| {
i32::try_from(d).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"slice_last_axis: dim",
"must fit in i32",
format_smolstr!("{d}"),
))
})
})
.collect::<Result<_>>()?;
let strides = vec![1_i32; n];
start[n - 1] = lo;
stop[n - 1] = hi;
ops::indexing::slice(x, &start, &stop, &strides)
}
struct Lstm {
wx_t: Array,
wh_t: Array,
bias: Array,
hidden_size: i32,
}
impl Lstm {
fn forward(
&self,
x: &Array,
hidden: Option<&Array>,
cell: Option<&Array>,
) -> Result<(Array, Array)> {
let pre = x.addmm(&self.bias, &self.wx_t, 1.0, 1.0)?;
let time_axis = pre.ndim() - 2;
let seq_len = i32::try_from(pre.shape()[time_axis]).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"silero_vad lstm: sequence length",
"must fit in i32",
format_smolstr!("{}", pre.shape()[time_axis]),
))
})?;
let time_axis_i = time_axis as i32;
let h = self.hidden_size;
let split_points = [h, 2 * h, 3 * h];
let mut hidden: Option<Array> = match hidden {
Some(h) => Some(h.try_clone()?),
None => None,
};
let mut cell: Option<Array> = match cell {
Some(c) => Some(c.try_clone()?),
None => None,
};
let mut all_hidden: Vec<Array> = Vec::with_capacity(seq_len as usize);
let mut all_cell: Vec<Array> = Vec::with_capacity(seq_len as usize);
for idx in 0..seq_len {
let mut ifgo = pre.take_axis(&idx0(idx)?, time_axis_i)?;
if let Some(prev_h) = &hidden {
ifgo = prev_h.addmm(&ifgo, &self.wh_t, 1.0, 1.0)?;
}
let parts = ops::shape::split_sections(&ifgo, &split_points, -1)?;
let i = ops::arithmetic::sigmoid(&parts[0])?;
let f = ops::arithmetic::sigmoid(&parts[1])?;
let g = ops::arithmetic::tanh(&parts[2])?;
let o = ops::arithmetic::sigmoid(&parts[3])?;
let new_cell = match &cell {
Some(prev_c) => f.multiply(prev_c)?.add(&i.multiply(&g)?)?,
None => i.multiply(&g)?,
};
let new_hidden = o.multiply(&ops::arithmetic::tanh(&new_cell)?)?;
cell = Some(new_cell.try_clone()?);
hidden = Some(new_hidden.try_clone()?);
all_cell.push(new_cell);
all_hidden.push(new_hidden);
}
let hidden_refs: Vec<&Array> = all_hidden.iter().collect();
let cell_refs: Vec<&Array> = all_cell.iter().collect();
let hidden_seq = ops::shape::stack_axis(&hidden_refs, -2)?;
let cell_seq = ops::shape::stack_axis(&cell_refs, -2)?;
Ok((hidden_seq, cell_seq))
}
}
pub struct SileroVadModel {
config: ModelConfig,
vad_16k: SileroVadBranch,
vad_8k: SileroVadBranch,
}
#[derive(Debug)]
pub struct SileroVadState {
state: Option<Array>,
context: Array,
sample_rate: u32,
}
impl SileroVadState {
#[inline(always)]
pub fn state(&self) -> Option<&Array> {
self.state.as_ref()
}
#[inline(always)]
pub fn context(&self) -> &Array {
&self.context
}
#[inline(always)]
pub const fn sample_rate(&self) -> u32 {
self.sample_rate
}
}
impl SileroVadModel {
pub fn new(config: ModelConfig, vad_16k: SileroVadBranch, vad_8k: SileroVadBranch) -> Self {
Self {
config,
vad_16k,
vad_8k,
}
}
#[inline(always)]
pub const fn config(&self) -> &ModelConfig {
&self.config
}
#[inline(always)]
pub const fn dtype(&self) -> Dtype {
self.config.dtype()
}
pub fn branch(&self, sample_rate: u32) -> Result<&SileroVadBranch> {
match sample_rate {
16_000 => Ok(&self.vad_16k),
8_000 => Ok(&self.vad_8k),
other => Err(Error::OutOfRange(OutOfRangePayload::new(
"silero_vad: sample_rate",
"Silero VAD supports 8000 Hz and 16000 Hz audio",
format_smolstr!("{other}"),
))),
}
}
pub fn forward(
&self,
x: &Array,
state: Option<&Array>,
sample_rate: u32,
) -> Result<(Array, Array)> {
let branch = self.branch(sample_rate)?;
let x = x.astype(self.dtype())?;
let state = match state {
Some(s) => Some(s.astype(self.dtype())?),
None => None,
};
branch.forward(&x, state.as_ref())
}
pub fn initial_state(&self, batch_size: i32, sample_rate: u32) -> Result<SileroVadState> {
let branch = self.branch(sample_rate)?;
let context =
Array::zeros::<f32>(&[batch_size, branch.config().context_size()])?.astype(self.dtype())?;
Ok(SileroVadState {
state: None,
context,
sample_rate,
})
}
pub fn feed(
&self,
chunk: &Array,
state: Option<SileroVadState>,
sample_rate: u32,
) -> Result<(Array, SileroVadState)> {
let ndim = chunk.ndim();
if ndim != 1 && ndim != 2 {
return Err(Error::RankMismatch(RankMismatchPayload::new(
"silero_vad feed: chunk must be rank-1 (T,) or rank-2 (B, T)",
ndim as u32,
chunk.shape().to_vec(),
)));
}
let branch = self.branch(sample_rate)?;
let chunk = chunk.astype(self.dtype())?;
let chunk = if chunk.ndim() == 1 {
ops::shape::expand_dims_axes(&chunk, &[0])?
} else {
chunk
};
let last_dim = *chunk.shape().last().unwrap_or(&0);
let chunk_width = i32::try_from(last_dim).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"silero_vad feed: chunk width",
"must fit in i32",
format_smolstr!("{last_dim}"),
))
})?;
if chunk_width != branch.config().chunk_size() {
return Err(Error::OutOfRange(OutOfRangePayload::new(
"silero_vad feed: chunk width",
"must equal the branch chunk_size for the sample rate",
format_smolstr!(
"expected={}, got={chunk_width}",
branch.config().chunk_size()
),
)));
}
let batch = i32::try_from(chunk.shape()[0]).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"silero_vad feed: batch size",
"must fit in i32",
format_smolstr!("{}", chunk.shape()[0]),
))
})?;
let state = match state {
Some(s) => s,
None => self.initial_state(batch, sample_rate)?,
};
if state.sample_rate != sample_rate {
return Err(Error::OutOfRange(OutOfRangePayload::new(
"silero_vad feed: state sample rate",
"streaming state sample rate must match the call",
format_smolstr!("state={}, call={sample_rate}", state.sample_rate),
)));
}
let window = ops::shape::concatenate(&[&state.context, &chunk], -1)?;
let (probability, lstm_state) = self.forward(&window, state.state.as_ref(), sample_rate)?;
let ctx = branch.config().context_size();
let new_context = slice_last_axis(&chunk, chunk_width - ctx, chunk_width)?;
Ok((
probability,
SileroVadState {
state: Some(lstm_state),
context: new_context,
sample_rate,
},
))
}
pub fn predict_proba(&self, audio: &Array, sample_rate: u32) -> Result<Array> {
let ndim = audio.ndim();
if ndim != 1 && ndim != 2 {
return Err(Error::RankMismatch(RankMismatchPayload::new(
"silero_vad predict_proba: audio must be rank-1 (T,) or rank-2 (B, T)",
ndim as u32,
audio.shape().to_vec(),
)));
}
let branch = self.branch(sample_rate)?;
let chunk_size = branch.config().chunk_size();
let context_size = branch.config().context_size();
let audio = audio.astype(self.dtype())?;
let original_ndim = audio.ndim();
let audio = if original_ndim == 1 {
ops::shape::expand_dims_axes(&audio, &[0])?
} else {
audio
};
let total = *audio.shape().last().unwrap_or(&0);
if total == 0 {
let batch = i32::try_from(audio.shape()[0]).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"silero_vad predict_proba: batch size",
"must fit in i32",
format_smolstr!("{}", audio.shape()[0]),
))
})?;
return if original_ndim == 1 {
Array::zeros::<f32>(&[0])?.astype(self.dtype())
} else {
Array::zeros::<f32>(&[batch, 0])?.astype(self.dtype())
};
}
let total_i = i32::try_from(total).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"silero_vad predict_proba: audio length",
"must fit in i32",
format_smolstr!("{total}"),
))
})?;
let batch = i32::try_from(audio.shape()[0]).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"silero_vad predict_proba: batch size",
"must fit in i32",
format_smolstr!("{}", audio.shape()[0]),
))
})?;
if batch == 0 {
let frames =
((i64::from(total_i) + i64::from(chunk_size) - 1) / i64::from(chunk_size)) as i32;
return Array::zeros::<f32>(&[0, frames])?.astype(self.dtype());
}
let pad = (chunk_size - total_i % chunk_size) % chunk_size;
let audio = if pad > 0 {
let zero = scalar_f32(0.0)?.astype(self.dtype())?;
ops::shape::pad(&audio, &[1], &[0], &[pad], &zero, c"constant")?
} else {
audio
};
let context = Array::zeros::<f32>(&[batch, context_size])?.astype(self.dtype())?;
let audio = ops::shape::concatenate(&[&context, &audio], -1)?;
let padded_len = i32::try_from(*audio.shape().last().unwrap_or(&0)).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"silero_vad predict_proba: padded length",
"must fit in i32",
"overflow",
))
})?;
let mut outputs: Vec<Array> = Vec::new();
let mut state: Option<Array> = None;
let mut pos = context_size;
let mut step = 0usize;
while pos < padded_len {
let window = slice_last_axis(&audio, pos - context_size, pos + chunk_size)?;
let (out, new_state) = self.forward(&window, state.as_ref(), sample_rate)?;
step += 1;
if step.is_multiple_of(EVAL_EVERY) {
crate::transforms::async_eval(&[&out, &new_state])?;
}
outputs.push(out);
state = Some(new_state);
pos += chunk_size;
}
if !outputs.len().is_multiple_of(EVAL_EVERY)
&& let (Some(last), Some(st)) = (outputs.last(), state.as_ref())
{
crate::transforms::async_eval(&[last, st])?;
}
let out_refs: Vec<&Array> = outputs.iter().collect();
let mut probabilities = ops::shape::concatenate(&out_refs, 1)?;
if original_ndim == 1 {
probabilities = probabilities.take_axis(&idx0(0)?, 0)?;
}
Ok(probabilities)
}
pub fn prepare_audio(&self, audio: &Array, sample_rate: u32) -> Result<(Array, u32)> {
if sample_rate == 0 {
return Err(Error::OutOfRange(OutOfRangePayload::new(
"silero_vad prepare_audio: sample_rate",
"must be > 0",
"0",
)));
}
let ndim = audio.ndim();
if ndim != 1 && ndim != 2 {
return Err(Error::RankMismatch(RankMismatchPayload::new(
"silero_vad prepare_audio: audio must be rank-1 (T,) or rank-2 (B, T) / (T, C)",
ndim as u32,
audio.shape(),
)));
}
if audio.shape().contains(&0) {
let target_sr = if sample_rate == 8_000 || sample_rate == 16_000 {
sample_rate
} else {
16_000
};
let shape = audio.shape();
let width = *shape.last().unwrap_or(&0);
if sample_rate != target_sr && shape.len() == 2 && width > 0 {
let resampled_width = usize::try_from(
(width as u64)
.checked_mul(u64::from(target_sr))
.ok_or_else(|| {
Error::ArithmeticOverflow(crate::error::ArithmeticOverflowPayload::with_operands(
"silero_vad prepare_audio: width * target_rate",
"u64",
[
("width", width as u64),
("target_rate", u64::from(target_sr)),
],
))
})?
/ u64::from(sample_rate),
)
.unwrap_or(usize::MAX);
let w = i32::try_from(resampled_width).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"silero_vad prepare_audio: resampled width",
"must fit in i32",
format_smolstr!("{resampled_width}"),
))
})?;
return Ok((
Array::zeros::<f32>(&[0, w])?.astype(audio.dtype()?)?,
target_sr,
));
}
return Ok((audio.try_clone()?, target_sr));
}
let downmixed = if ndim == 2 {
let shape = audio.shape();
let rows = shape[0];
let cols = shape[1];
if cols <= 8 && 8 < rows {
ops::reduction::mean_axes(audio, &[-1], false)?
} else {
audio.try_clone()?
}
} else {
audio.try_clone()?
};
let target_sr = if sample_rate == 8_000 || sample_rate == 16_000 {
sample_rate
} else {
16_000
};
if sample_rate == target_sr {
return Ok((downmixed, sample_rate));
}
Ok((
resample_audio_array(&downmixed, sample_rate, target_sr)?,
target_sr,
))
}
pub fn predict(&self, audio: &Array, sample_rate: u32) -> Result<Array> {
let (audio, sr) = self.prepare_audio(audio, sample_rate)?;
self.predict_proba(&audio, sr)
}
pub fn reset_state(&self, batch_size: i32, sample_rate: u32) -> Result<SileroVadState> {
self.initial_state(batch_size, sample_rate)
}
pub fn get_speech_timestamps(
&self,
audio: &Array,
sample_rate: u32,
options: SpeechTimestampOptions,
) -> Result<Vec<SpeechSegment>> {
let (audio, sr) = self.prepare_audio(audio, sample_rate)?;
let audio_len = *audio.shape().last().unwrap_or(&0) as i64;
let mut probabilities = self.predict_proba(&audio, sr)?;
probabilities.eval()?;
let probs_vec: Vec<f32> = if probabilities.ndim() == 2 {
if probabilities.shape()[0] == 0 {
Vec::new()
} else {
let mut row = probabilities.take_axis(&idx0(0)?, 0)?;
row.eval()?;
row.astype(Dtype::F32)?.to_vec::<f32>()?
}
} else {
probabilities.astype(Dtype::F32)?.to_vec::<f32>()?
};
let cfg = &self.config;
let threshold = options.threshold.unwrap_or(cfg.threshold());
let min_speech_duration_ms = options
.min_speech_duration_ms
.unwrap_or(cfg.min_speech_duration_ms());
let min_silence_duration_ms = options
.min_silence_duration_ms
.unwrap_or(cfg.min_silence_duration_ms());
let speech_pad_ms = options.speech_pad_ms.unwrap_or(cfg.speech_pad_ms());
if !threshold.is_finite() {
return Err(Error::OutOfRange(OutOfRangePayload::new(
"silero_vad get_speech_timestamps: threshold",
"must be finite",
format_smolstr!("{threshold}"),
)));
}
for (field, v) in [
("min_speech_duration_ms", min_speech_duration_ms),
("min_silence_duration_ms", min_silence_duration_ms),
("speech_pad_ms", speech_pad_ms),
] {
if v < 0 {
return Err(Error::OutOfRange(OutOfRangePayload::new(
"silero_vad get_speech_timestamps: negative duration/padding override",
"must be >= 0",
format_smolstr!("{field}={v}"),
)));
}
}
Ok(probs_to_timestamps(
&probs_vec,
audio_len,
sr,
threshold,
min_speech_duration_ms,
min_silence_duration_ms,
speech_pad_ms,
))
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct SpeechTimestampOptions {
pub threshold: Option<f64>,
pub min_speech_duration_ms: Option<i32>,
pub min_silence_duration_ms: Option<i32>,
pub speech_pad_ms: Option<i32>,
}
fn resample_audio_array(audio: &Array, from: u32, to: u32) -> Result<Array> {
let shape = audio.shape();
if shape.contains(&0) {
return audio.try_clone();
}
let resample_1d = |samples: &[f32], leading_batch: bool| -> Result<Array> {
let out = crate::audio::io::resample_linear(samples, from, to)?;
let n = i32::try_from(out.len()).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"silero_vad prepare_audio: resampled length",
"must fit in i32",
format_smolstr!("{}", out.len()),
))
})?;
if leading_batch {
Array::from_slice::<f32>(&out, &[1, n])
} else {
Array::from_slice::<f32>(&out, &[n])
}
};
match shape.len() {
1 => {
let samples = audio.astype(Dtype::F32)?.to_vec::<f32>()?;
resample_1d(&samples, false)
}
2 => {
let rows = shape[0];
let f32_audio = audio.astype(Dtype::F32)?;
let mut resampled_rows: Vec<Array> = Vec::with_capacity(rows);
for r in 0..i32::try_from(rows).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"silero_vad prepare_audio: batch size",
"must fit in i32",
format_smolstr!("{rows}"),
))
})? {
let samples = f32_audio.take_axis(&idx0(r)?, 0)?.to_vec::<f32>()?;
resampled_rows.push(resample_1d(&samples, true)?);
}
let refs: Vec<&Array> = resampled_rows.iter().collect();
ops::shape::concatenate(&refs, 0)
}
_ => Err(Error::RankMismatch(RankMismatchPayload::new(
"silero_vad prepare_audio: resample rank",
shape.len() as u32,
shape,
))),
}
}
#[derive(Debug, Clone, Copy)]
struct Speech {
start: i64,
end: i64,
}
#[allow(clippy::too_many_arguments)]
pub fn probs_to_timestamps(
probs: &[f32],
audio_len: i64,
sample_rate: u32,
threshold: f64,
min_speech_duration_ms: i32,
min_silence_duration_ms: i32,
speech_pad_ms: i32,
) -> Vec<SpeechSegment> {
let sr = sample_rate as f64;
let chunk_size: i64 = if sample_rate == 16_000 { 512 } else { 256 };
let min_speech_samples = sr * f64::from(min_speech_duration_ms) / 1000.0;
let min_silence_samples = sr * f64::from(min_silence_duration_ms) / 1000.0;
let speech_pad_samples = (sr * f64::from(speech_pad_ms) / 1000.0) as i64;
let neg_threshold = (threshold - 0.15).max(0.01);
let mut speeches: Vec<Speech> = Vec::new();
let mut triggered = false;
let mut current_start: i64 = 0;
let mut temp_end: i64 = 0;
for (idx, &prob) in probs.iter().enumerate() {
let prob = f64::from(prob);
let chunk_start = idx as i64 * chunk_size;
if prob >= threshold && !triggered {
triggered = true;
current_start = chunk_start;
temp_end = 0;
continue;
}
if triggered && prob >= threshold {
temp_end = 0;
continue;
}
if triggered && prob < neg_threshold {
if temp_end == 0 {
temp_end = chunk_start;
}
if (chunk_start - temp_end) as f64 >= min_silence_samples {
if (temp_end - current_start) as f64 >= min_speech_samples {
speeches.push(Speech {
start: current_start,
end: temp_end,
});
}
triggered = false;
temp_end = 0;
}
}
}
if triggered {
let end = audio_len.min(probs.len() as i64 * chunk_size);
if (end - current_start) as f64 >= min_speech_samples {
speeches.push(Speech {
start: current_start,
end,
});
}
}
let mut padded: Vec<Speech> = Vec::new();
for speech in &speeches {
let start = (speech.start - speech_pad_samples).max(0);
let end = audio_len.min(speech.end + speech_pad_samples);
if let Some(last) = padded.last_mut()
&& start <= last.end
{
last.end = last.end.max(end);
} else {
padded.push(Speech { start, end });
}
}
padded
.into_iter()
.map(|s| SpeechSegment::new(s.start.max(0) as u64, s.end.max(0) as u64))
.collect()
}
pub fn sanitize(weights: HashMap<String, Array>) -> HashMap<String, Array> {
weights
.into_iter()
.filter(|(k, _)| !k.starts_with("val_"))
.collect()
}
#[allow(clippy::too_many_arguments)]
pub(super) fn build_branch(
config: BranchConfig,
stft_conv_weight: Array,
conv1: (Array, Array),
conv2: (Array, Array),
conv3: (Array, Array),
conv4: (Array, Array),
lstm_wx: &Array,
lstm_wh: &Array,
lstm_bias: Array,
final_conv_weight: Array,
final_conv_bias: Array,
) -> Result<SileroVadBranch> {
let wx_shape = lstm_wx.shape();
let wh_shape = lstm_wh.shape();
let bias_shape = lstm_bias.shape();
if wx_shape.len() != 2 || wh_shape.len() != 2 || bias_shape.len() != 1 {
return Err(Error::RankMismatch(RankMismatchPayload::new(
"silero_vad lstm weights: expected Wx (4H, D), Wh (4H, H), bias (4H,)",
wx_shape.len() as u32,
wx_shape,
)));
}
let four_h = wx_shape[0];
if four_h == 0 || !four_h.is_multiple_of(4) {
return Err(Error::OutOfRange(OutOfRangePayload::new(
"silero_vad lstm Wx leading dim",
"must be a positive multiple of 4 (the 4*hidden gate stack)",
format_smolstr!("{four_h}"),
)));
}
let hidden = four_h / 4;
if wh_shape[0] != four_h || wh_shape[1] != hidden || bias_shape[0] != four_h {
return Err(Error::OutOfRange(OutOfRangePayload::new(
"silero_vad lstm weights: Wh / bias inconsistent with Wx",
"Wh must be (4H, H) and bias (4H,) for H = Wx.shape[0] / 4",
format_smolstr!("Wx={wx_shape:?}, Wh={wh_shape:?}, bias={bias_shape:?}"),
)));
}
let hidden_size = i32::try_from(hidden).map_err(|_| {
Error::OutOfRange(OutOfRangePayload::new(
"silero_vad lstm hidden size",
"must fit in i32",
format_smolstr!("{hidden}"),
))
})?;
let wx_t = ops::shape::swapaxes(lstm_wx, -2, -1)?;
let wh_t = ops::shape::swapaxes(lstm_wh, -2, -1)?;
Ok(SileroVadBranch {
config,
stft_conv_weight,
conv1: ConvBlock {
weight: conv1.0,
bias: conv1.1,
stride: 1,
padding: 1,
},
conv2: ConvBlock {
weight: conv2.0,
bias: conv2.1,
stride: 2,
padding: 1,
},
conv3: ConvBlock {
weight: conv3.0,
bias: conv3.1,
stride: 2,
padding: 1,
},
conv4: ConvBlock {
weight: conv4.0,
bias: conv4.1,
stride: 1,
padding: 1,
},
lstm: Lstm {
wx_t,
wh_t,
bias: lstm_bias,
hidden_size,
},
final_conv_weight,
final_conv_bias,
})
}