From 6769ad376ef448590b241248a65f923e3b36f66f Mon Sep 17 00:00:00 2001 From: jhqxxx <18280426169@163.com> Date: Fri, 17 Apr 2026 18:23:26 +0800 Subject: [PATCH] vad add speech cache --- src/models/common/modules.rs | 4 +-- src/models/fire_red_vad/vad.rs | 51 ++++++++++++++++++++++++++++---- src/models/qwen3_asr/generate.rs | 3 -- 3 files changed, 47 insertions(+), 11 deletions(-) diff --git a/src/models/common/modules.rs b/src/models/common/modules.rs index cc0ef8a..bc0fb8b 100644 --- a/src/models/common/modules.rs +++ b/src/models/common/modules.rs @@ -15,10 +15,10 @@ use crate::{ #[derive(Debug)] pub struct VadFrameResult { pub is_speech: bool, - pub is_speech_start: bool, + // pub is_speech_start: bool, pub is_i16: bool, pub orig_audio: Option, - pub kaldi_audio: Option, + // pub kaldi_audio: Option, pub model_name: String, pub mode: String, } diff --git a/src/models/fire_red_vad/vad.rs b/src/models/fire_red_vad/vad.rs index 17a9c17..9889881 100644 --- a/src/models/fire_red_vad/vad.rs +++ b/src/models/fire_red_vad/vad.rs @@ -35,6 +35,7 @@ pub struct FireRedVad { cfg: FireRedVadConfig, caches: Option>, frame_length_sample: usize, + speech_cache: Vec, } impl FireRedVad { @@ -76,6 +77,7 @@ impl FireRedVad { cfg, caches: None, frame_length_sample: 400, + speech_cache: vec![], }) } @@ -98,18 +100,55 @@ impl FireRedVad { .process_thresh(&probs)? .to_dtype(DType::U32)?; let preds_sum = binary_preds.sum_all()?.to_scalar::()?; - if preds_sum as f32 > probs.dim(0)? as f32 * self.cfg.speech_threshold { + let probs_len = probs.dim(0)?; + // 输入数据中 is_speech > 0.1, 认为这帧数据可用 + let final_data = if preds_sum as f32 > probs_len as f32 * 0.1 { + // 通过最后10个数据,判断说话是否结束 + let last_10_preds_sum = binary_preds + .narrow(0, probs_len - 10, 10)? + .sum_all()? + .to_scalar::()?; + // 10个数据中,至少8个是 speech, 认为说话没有结束,缓存数据,等待下一帧 + if last_10_preds_sum >= 8 { + self.speech_cache.push(audio_frame.clone()); + None + } else { + // 否则认为此次说话结束,结合缓存数据,一起返回 + let data = if self.speech_cache.is_empty() { + // 缓存数据为空,直接返回 + audio_frame.clone() + } else { + // 缓存数据不为空,cat数据 + self.speech_cache.push(audio_frame.clone()); + let audio_frame = Tensor::cat(&self.speech_cache, 0)?; + self.speech_cache = vec![]; // 清空缓存 + audio_frame + }; + Some(data) + } + } else { + // 认为这帧数据不可用,是否有缓存数据 + if self.speech_cache.is_empty() { + // 没有返回None + None + } else { + // 有缓存,返回缓存数据 + let data = Some(Tensor::cat(&self.speech_cache, 0)?); + self.speech_cache.clear(); // 清空缓存 + data + } + }; + + if final_data.is_none() { + Ok(None) + } else { Ok(Some(VadFrameResult { is_speech: true, is_i16: true, - is_speech_start: true, // TODO: is start speech, asr to clear cache - orig_audio: Some(audio_frame.clone()), - kaldi_audio: Some(feats), + orig_audio: final_data, model_name: self.model_name.clone(), mode: "speech".to_string(), })) - } else { - Ok(None) } } diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs index c6befe3..18810f5 100644 --- a/src/models/qwen3_asr/generate.rs +++ b/src/models/qwen3_asr/generate.rs @@ -89,9 +89,6 @@ impl<'a> Qwen3AsrGenerateModel<'a> { if !vad_res.is_speech || vad_res.orig_audio.is_none() { return Ok(AsrResult::init_empty()); } - // if vad_res.is_speech_start { - // self.qwen3_asr.clear_kv_cache(); - // } let audio_data = self.processor .process_vad_res(&self.default_template, vad_res, &self.tokenizer)?;