fix FireRedVad detect_frame bug

This commit is contained in:
jhqxxx
2026-04-20 15:15:13 +08:00
parent deb8027a82
commit 0017bcd924
4 changed files with 33 additions and 31 deletions
+5 -4
View File
@@ -103,13 +103,14 @@ impl FireRedVad {
let probs_len = probs.dim(0)?; let probs_len = probs.dim(0)?;
// 输入数据中 is_speech > 0.1, 认为这帧数据可用 // 输入数据中 is_speech > 0.1, 认为这帧数据可用
let final_data = if preds_sum as f32 > probs_len as f32 * 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 let last_10_preds_sum = binary_preds
.narrow(0, probs_len - 10, 10)? .narrow(0, probs_len - select_len, select_len)?
.sum_all()? .sum_all()?
.to_scalar::<u32>()?; .to_scalar::<u32>()?;
// 10个数据中,至少8个是 speech, 认为说话没有结束,缓存数据,等待下一帧 // 选中数据中,至少0.8个是 speech, 认为说话没有结束,缓存数据,等待下一帧
if last_10_preds_sum >= 8 { if last_10_preds_sum >= (select_len as f32 * 0.8).ceil() as u32 {
self.speech_cache.push(audio_frame.clone()); self.speech_cache.push(audio_frame.clone());
None None
} else { } else {
+1 -1
View File
@@ -388,7 +388,7 @@ pub fn load_audio_use_hound(audio_path: PathBuf, device: &Device) -> Result<(Ten
.collect::<Result<Vec<_>, _>>()?, .collect::<Result<Vec<_>, _>>()?,
24 => reader 24 => reader
.samples::<i32>() .samples::<i32>()
.map(|s| s.map(|sample| sample as f32 / 8388607.0)) .map(|s| s.map(|sample| sample as f32 / i32::MAX as f32))
.collect::<Result<Vec<_>, _>>()?, .collect::<Result<Vec<_>, _>>()?,
_ => { _ => {
return Err(anyhow::anyhow!( return Err(anyhow::anyhow!(
+1 -1
View File
@@ -23,7 +23,7 @@ fn stream_vad() -> Result<()> {
let device = aha::Device::Cpu; let device = aha::Device::Cpu;
let save_dir = let save_dir =
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get 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 vad = FireRedVad::init(&model_path, Some(&device), None)?;
let res = vad.detect_file(audio_path)?; let res = vad.detect_file(audio_path)?;
+26 -25
View File
@@ -60,7 +60,7 @@ fn voxcpm_use_message_generate() -> Result<()> {
#[test] #[test]
fn voxcpm_generate() -> Result<()> { 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 = let save_dir =
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get 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); let model_path = format!("{}/OpenBMB/VoxCPM-0.5B/", save_dir);
@@ -70,38 +70,39 @@ fn voxcpm_generate() -> Result<()> {
let i_duration = i_start.elapsed(); let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration); 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.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?;
let generate = voxcpm_generate.inference( // let generate = voxcpm_generate.inference(
"VoxCPM is an innovative end-to-end TTS model from ModelBest, designed to generate highly realistic speech.".to_string(), // "VoxCPM is an innovative end-to-end TTS model from ModelBest, designed to generate highly realistic speech.".to_string(),
Some("啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string()), // Some("啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string()),
Some("file://./assets/audio/voice_01.wav".to_string()), // Some("file://./assets/audio/voice_01.wav".to_string()),
// Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()), // // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
// Some("file://./assets/audio/voice_05.wav".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(),
// 2, // 2,
// 100, // 100,
// 10, // 10,
// 2.0, // 2.0,
// false, // // false,
// 6.0, // 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(); let i_duration = i_start.elapsed();
println!("Time elapsed in generate is: {:?}", i_duration); println!("Time elapsed in generate is: {:?}", i_duration);
save_wav(&generate, "voxcpm.wav", 16000)?; save_wav(&generate, "voxcpm.wav", 16000)?;