From 0017bcd924c78b18d438acd79841a651b2688948 Mon Sep 17 00:00:00 2001 From: jhqxxx <18280426169@163.com> Date: Mon, 20 Apr 2026 15:15:13 +0800 Subject: [PATCH] fix FireRedVad detect_frame bug --- src/models/fire_red_vad/vad.rs | 9 +++--- src/utils/audio_utils.rs | 2 +- tests/test_fire_red_vad.rs | 2 +- tests/test_voxcpm.rs | 51 +++++++++++++++++----------------- 4 files changed, 33 insertions(+), 31 deletions(-) diff --git a/src/models/fire_red_vad/vad.rs b/src/models/fire_red_vad/vad.rs index 9889881..6631df8 100644 --- a/src/models/fire_red_vad/vad.rs +++ b/src/models/fire_red_vad/vad.rs @@ -103,13 +103,14 @@ impl FireRedVad { 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个数据,判断说话是否结束 + // 通过最后10个数据,判断说话是否结束,如果数据长度小于10就取整个长度 + let select_len = if probs_len > 10 { 10 } else { probs_len }; let last_10_preds_sum = binary_preds - .narrow(0, probs_len - 10, 10)? + .narrow(0, probs_len - select_len, select_len)? .sum_all()? .to_scalar::()?; - // 10个数据中,至少8个是 speech, 认为说话没有结束,缓存数据,等待下一帧 - if last_10_preds_sum >= 8 { + // 选中数据中,至少0.8个是 speech, 认为说话没有结束,缓存数据,等待下一帧 + if last_10_preds_sum >= (select_len as f32 * 0.8).ceil() as u32 { self.speech_cache.push(audio_frame.clone()); None } else { diff --git a/src/utils/audio_utils.rs b/src/utils/audio_utils.rs index f104831..6ab5eef 100644 --- a/src/utils/audio_utils.rs +++ b/src/utils/audio_utils.rs @@ -388,7 +388,7 @@ pub fn load_audio_use_hound(audio_path: PathBuf, device: &Device) -> Result<(Ten .collect::, _>>()?, 24 => reader .samples::() - .map(|s| s.map(|sample| sample as f32 / 8388607.0)) + .map(|s| s.map(|sample| sample as f32 / i32::MAX as f32)) .collect::, _>>()?, _ => { return Err(anyhow::anyhow!( diff --git a/tests/test_fire_red_vad.rs b/tests/test_fire_red_vad.rs index c06d69c..e0a2ef9 100644 --- a/tests/test_fire_red_vad.rs +++ b/tests/test_fire_red_vad.rs @@ -23,7 +23,7 @@ fn stream_vad() -> Result<()> { let device = aha::Device::Cpu; let save_dir = aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; - let model_path = format!("{}/xukaituo/FireRedVAD/Stream-VAD/", save_dir); + let model_path = format!("{}/jiangjiangaha/FireRedVAD-Stream-VAD/", save_dir); let vad = FireRedVad::init(&model_path, Some(&device), None)?; let res = vad.detect_file(audio_path)?; diff --git a/tests/test_voxcpm.rs b/tests/test_voxcpm.rs index 3866a7e..76b09ed 100644 --- a/tests/test_voxcpm.rs +++ b/tests/test_voxcpm.rs @@ -60,7 +60,7 @@ fn voxcpm_use_message_generate() -> Result<()> { #[test] fn voxcpm_generate() -> Result<()> { - // RUST_BACKTRACE=1 cargo test -F cuda voxcpm_generate -r -- --nocapture + // RUST_BACKTRACE=1 cargo test -F cuda --test test_voxcpm voxcpm_generate -r -- --nocapture let save_dir = aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; let model_path = format!("{}/OpenBMB/VoxCPM-0.5B/", save_dir); @@ -70,38 +70,39 @@ fn voxcpm_generate() -> Result<()> { let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); + // let i_start = Instant::now(); // let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?; - let generate = voxcpm_generate.inference( - "VoxCPM is an innovative end-to-end TTS model from ModelBest, designed to generate highly realistic speech.".to_string(), - Some("啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string()), - Some("file://./assets/audio/voice_01.wav".to_string()), - // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()), - // Some("file://./assets/audio/voice_05.wav".to_string()), - 2, - 100, - 10, - 2.0, - // false, - 6.0, - )?; - - // 创建prompt_cache - // let _ = voxcpm_generate.build_prompt_cache( - // "啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(), - // "file://./assets/audio/voice_01.wav".to_string(), - // )?; - // // 使用prompt_cache生成语音 - // let generate = voxcpm_generate.generate_use_prompt_cache( - // "太阳当空照,花儿对我笑,小鸟说早早早".to_string(), + // let generate = voxcpm_generate.inference( + // "VoxCPM is an innovative end-to-end TTS model from ModelBest, designed to generate highly realistic speech.".to_string(), + // Some("啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string()), + // Some("file://./assets/audio/voice_01.wav".to_string()), + // // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()), + // // Some("file://./assets/audio/voice_05.wav".to_string()), // 2, // 100, // 10, // 2.0, - // false, + // // false, // 6.0, // )?; + // 创建prompt_cache + let _ = voxcpm_generate.build_prompt_cache( + "啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(), + "file://./assets/audio/voice_01.wav".to_string(), + )?; + // 使用prompt_cache生成语音 + let i_start = Instant::now(); + let generate = voxcpm_generate.generate_use_prompt_cache( + "太阳当空照,花儿对我笑,小鸟说早早早".to_string(), + 2, + 100, + 10, + 2.0, + false, + 6.0, + )?; + let i_duration = i_start.elapsed(); println!("Time elapsed in generate is: {:?}", i_duration); save_wav(&generate, "voxcpm.wav", 16000)?;