update fire_red_vad

This commit is contained in:
jhqxxx
2026-04-16 19:47:28 +08:00
parent 76f9c202a5
commit 26fedf6bef
2 changed files with 69 additions and 12 deletions
+35 -10
View File
@@ -9,7 +9,8 @@ use crate::{
processor::{AudioFeat, VadPostprocessor}, processor::{AudioFeat, VadPostprocessor},
}, },
utils::{ utils::{
audio_utils::load_audio_with_resample_from_bytes, find_type_files, get_device, audio_utils::{resample_audio_from_bytes, resample_audio_from_vec_f32},
find_type_files, get_device,
tensor_utils::split_tensor_with_size, tensor_utils::split_tensor_with_size,
}, },
}; };
@@ -84,13 +85,7 @@ impl FireRedVad {
}) })
} }
pub fn detect_frame(&mut self, audio_bytes: Vec<u8>) -> Result<Option<VadFrameResult>> { pub fn detect_frame(&mut self, audio_frame: &Tensor) -> Result<Option<VadFrameResult>> {
if !self.model_name.to_lowercase().contains("stream") {
return Err(anyhow!("only stream model support detect_frame"));
}
let audio_frame =
load_audio_with_resample_from_bytes(audio_bytes, &self.device, Some(16000), true)?
.squeeze(0)?;
if audio_frame.dim(0)? < self.frame_length_sample { if audio_frame.dim(0)? < self.frame_length_sample {
return Err(anyhow!( return Err(anyhow!(
"Expected {} samples, got {}", "Expected {} samples, got {}",
@@ -98,7 +93,7 @@ impl FireRedVad {
audio_frame.dim(0)? audio_frame.dim(0)?
)); ));
} }
let feats = self.audio_feat.extract(&audio_frame)?; let feats = self.audio_feat.extract(audio_frame)?;
let (probs, caches) = self let (probs, caches) = self
.vad_model .vad_model
.forward(&feats.unsqueeze(0)?, self.caches.as_ref())?; .forward(&feats.unsqueeze(0)?, self.caches.as_ref())?;
@@ -112,7 +107,7 @@ impl FireRedVad {
if preds_sum as f32 > probs.dim(0)? as f32 * self.cfg.speech_threshold { if preds_sum as f32 > probs.dim(0)? as f32 * self.cfg.speech_threshold {
Ok(Some(VadFrameResult { Ok(Some(VadFrameResult {
is_speech: true, is_speech: true,
orig_audio: Some(audio_frame), orig_audio: Some(audio_frame.clone()),
kaldi_audio: Some(feats), kaldi_audio: Some(feats),
model_name: self.model_name.clone(), model_name: self.model_name.clone(),
mode: "speech".to_string(), mode: "speech".to_string(),
@@ -122,6 +117,36 @@ impl FireRedVad {
} }
} }
pub fn detect_frame_f32(
&mut self,
audio_vec_f32: Vec<f32>,
channels: usize,
orig_sr: Option<usize>,
) -> Result<Option<VadFrameResult>> {
if !self.model_name.to_lowercase().contains("stream") {
return Err(anyhow!("only stream model support detect_frame"));
}
let audio_frame = resample_audio_from_vec_f32(
audio_vec_f32,
&self.device,
channels,
orig_sr,
Some(16000),
true,
)?
.squeeze(0)?;
self.detect_frame(&audio_frame)
}
pub fn detect_frame_bytes(&mut self, audio_bytes: Vec<u8>) -> Result<Option<VadFrameResult>> {
if !self.model_name.to_lowercase().contains("stream") {
return Err(anyhow!("only stream model support detect_frame"));
}
let audio_frame =
resample_audio_from_bytes(audio_bytes, &self.device, Some(16000), true)?.squeeze(0)?;
self.detect_frame(&audio_frame)
}
pub fn detect_file(&self, audio_path: &str) -> Result<VadResult> { pub fn detect_file(&self, audio_path: &str) -> Result<VadResult> {
let (feats, dur) = self.audio_feat.extract_file(audio_path, &self.device)?; let (feats, dur) = self.audio_feat.extract_file(audio_path, &self.device)?;
let probs = if feats.dim(0)? <= self.cfg.chunk_max_frame { let probs = if feats.dim(0)? <= self.cfg.chunk_max_frame {
+34 -2
View File
@@ -587,7 +587,39 @@ pub fn load_audio(path: &str, device: &Device) -> Result<(Tensor, usize)> {
load_audio_use_symphonia(audio_vec, false, device) load_audio_use_symphonia(audio_vec, false, device)
} }
pub fn load_audio_with_resample_from_bytes( pub fn resample_audio_from_vec_f32(
audio_vec: Vec<f32>,
device: &Device,
channels: usize,
orig_sr: Option<usize>,
target_sample_rate: Option<usize>,
is_i16: bool,
) -> Result<Tensor> {
let frame_len = audio_vec.len() / channels;
let audio = Tensor::new(&audio_vec[0..frame_len * channels], device)?;
let mut audio = if channels > 1 {
audio
.reshape((frame_len, channels))?
.mean_keepdim(1)?
.transpose(0, 1)?
.contiguous()?
} else {
audio.unsqueeze(0)?
};
if let Some(target_sample_rate) = target_sample_rate
&& let Some(sr) = orig_sr
&& target_sample_rate != sr
{
audio = resample_simple(&audio, sr as i64, target_sample_rate as i64)?;
}
if is_i16 {
audio = audio.affine(32768.0, 0.0)?;
}
Ok(audio)
}
pub fn resample_audio_from_bytes(
audio_vec: Vec<u8>, audio_vec: Vec<u8>,
device: &Device, device: &Device,
target_sample_rate: Option<usize>, target_sample_rate: Option<usize>,
@@ -612,7 +644,7 @@ pub fn load_audio_with_resample(
// let audio_path = get_audio_path(path)?; // let audio_path = get_audio_path(path)?;
// let (mut audio, sr) = load_audio_use_hound(audio_path, device)?; // let (mut audio, sr) = load_audio_use_hound(audio_path, device)?;
let audio_vec = get_audio_bytes_vec(path)?; let audio_vec = get_audio_bytes_vec(path)?;
load_audio_with_resample_from_bytes(audio_vec, device, target_sample_rate, is_i16) resample_audio_from_bytes(audio_vec, device, target_sample_rate, is_i16)
} }
pub fn save_wav(audio: &Tensor, save_path: &str, sample_rate: u32) -> Result<()> { pub fn save_wav(audio: &Tensor, save_path: &str, sample_rate: u32) -> Result<()> {