use std::{
collections::HashMap,
path::{Path, PathBuf},
};
use smol_str::format_smolstr;
use crate::{
array::Array,
audio::vad::{
load::VadModel,
models::silero_vad::{
config::{BranchConfig, ModelConfig},
model::{SileroVadBranch, SileroVadModel, build_branch, sanitize},
},
output::{SpeechSegment, VadOutput},
},
error::{
Error, FileIoPayload, FileOp, LayerKeyedPayload, MissingKeyPayload, OutOfRangePayload, Result,
},
};
impl VadModel for SileroVadModel {
fn generate(&self, audio: &Array, sample_rate: u32) -> Result<VadOutput> {
let (audio, sample_rate) = self.prepare_audio(audio, sample_rate)?;
let audio_len = *audio.shape().last().unwrap_or(&0) as i64;
let mut probabilities = self.predict_proba(&audio, sample_rate)?;
probabilities.eval()?;
let probs_vec: Vec<f32> = if probabilities.ndim() == 2 {
if probabilities.shape()[0] == 0 {
Vec::new()
} else {
probabilities
.take_axis(&Array::from_slice::<i32>(&[0], &[0i32; 0])?, 0)?
.astype(crate::dtype::Dtype::F32)?
.to_vec::<f32>()?
}
} else {
probabilities
.try_clone()?
.astype(crate::dtype::Dtype::F32)?
.to_vec::<f32>()?
};
let cfg = self.config();
let timestamps: Vec<SpeechSegment> =
crate::audio::vad::models::silero_vad::model::probs_to_timestamps(
&probs_vec,
audio_len,
sample_rate,
cfg.threshold(),
cfg.min_speech_duration_ms(),
cfg.min_silence_duration_ms(),
cfg.speech_pad_ms(),
);
Ok(VadOutput {
timestamps,
probabilities,
sample_rate,
})
}
}
impl SileroVadModel {
pub fn from_weights(config: ModelConfig, mut weights: HashMap<String, Array>) -> Result<Self> {
if has_relevant_scales(&weights) {
return Err(Error::OutOfRange(OutOfRangePayload::new(
"silero_vad: quantized checkpoint",
"Silero VAD is dense-only; a quantized checkpoint (with .scales tensors) is unsupported",
"quantized",
)));
}
let vad_16k = build_branch_from_weights(*config.branch_16k(), &mut weights, "vad_16k")?;
let vad_8k = build_branch_from_weights(*config.branch_8k(), &mut weights, "vad_8k")?;
Ok(Self::new(config, vad_16k, vad_8k))
}
pub fn load(path: &str) -> Result<Self> {
let dir = crate::audio::load::get_model_path(path)?;
let config_json = crate::audio::load::load_config(&dir)?;
let config = ModelConfig::from_json(&config_json)?;
let raw = load_all_safetensors(&dir)?;
let weights = sanitize(raw);
Self::from_weights(config, weights)
}
}
fn build_branch_from_weights(
config: BranchConfig,
weights: &mut HashMap<String, Array>,
branch: &str,
) -> Result<SileroVadBranch> {
let stft = take_weight(weights, branch, "stft_conv.weight")?;
let conv1 = (
take_weight(weights, branch, "conv1.weight")?,
take_weight(weights, branch, "conv1.bias")?,
);
let conv2 = (
take_weight(weights, branch, "conv2.weight")?,
take_weight(weights, branch, "conv2.bias")?,
);
let conv3 = (
take_weight(weights, branch, "conv3.weight")?,
take_weight(weights, branch, "conv3.bias")?,
);
let conv4 = (
take_weight(weights, branch, "conv4.weight")?,
take_weight(weights, branch, "conv4.bias")?,
);
let lstm_wx = take_weight(weights, branch, "lstm.Wx")?;
let lstm_wh = take_weight(weights, branch, "lstm.Wh")?;
let lstm_bias = take_weight(weights, branch, "lstm.bias")?;
let final_w = take_weight(weights, branch, "final_conv.weight")?;
let final_b = take_weight(weights, branch, "final_conv.bias")?;
build_branch(
config, stft, conv1, conv2, conv3, conv4, &lstm_wx, &lstm_wh, lstm_bias, final_w, final_b,
)
}
fn take_weight(weights: &mut HashMap<String, Array>, branch: &str, suffix: &str) -> Result<Array> {
let key = format!("{branch}.{suffix}");
weights
.remove(&key)
.ok_or_else(|| Error::MissingKey(MissingKeyPayload::new("silero_vad: missing weight", key)))
}
pub fn has_relevant_scales(weights: &HashMap<String, Array>) -> bool {
weights.keys().any(|k| k.ends_with(".scales"))
}
fn load_all_safetensors(dir: &Path) -> Result<HashMap<String, Array>> {
let entries = std::fs::read_dir(dir).map_err(|e| {
Error::FileIo(FileIoPayload::new(
"silero_vad load: read model directory",
FileOp::Read,
dir.to_path_buf(),
e,
))
})?;
let mut files: Vec<PathBuf> = entries
.map(|entry| {
entry.map(|e| e.path()).map_err(|e| {
Error::FileIo(FileIoPayload::new(
"silero_vad load: read model directory entry",
FileOp::Read,
dir.to_path_buf(),
e,
))
})
})
.collect::<Result<Vec<_>>>()?;
files.retain(|p| p.extension().is_some_and(|ext| ext == "safetensors"));
files.sort();
if files.is_empty() {
return Err(Error::MissingKey(MissingKeyPayload::new(
"silero_vad load: no *.safetensors in model directory",
format_smolstr!("{}", dir.display()),
)));
}
let mut all = HashMap::new();
for f in &files {
let shard = crate::io::load_safetensors(f)?;
for (key, value) in shard {
crate::model_validation::insert_unique(
&mut all,
key,
value,
"silero_vad load: duplicate tensor key across shards",
)
.map_err(|e| match e {
Error::KeyCollision(_) => {
Error::LayerKeyed(LayerKeyedPayload::new(f.to_string_lossy().into_owned(), e))
}
other => other,
})?;
}
}
Ok(all)
}
pub fn load(path: &str) -> Result<Box<dyn VadModel>> {
Ok(Box::new(SileroVadModel::load(path)?))
}