fix FireRedVad detect_frame bug
This commit is contained in:
@@ -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 {
|
||||||
|
|||||||
@@ -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!(
|
||||||
|
|||||||
@@ -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
@@ -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)?;
|
||||||
|
|||||||
Reference in New Issue
Block a user