diff --git a/README.md b/README.md index 2d71328..71caa29 100644 --- a/README.md +++ b/README.md @@ -33,12 +33,15 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an | **Vision** | Qwen2.5-VL, Qwen3-VL, Qwen3.5,
LFM2.5-VL, LFM2-VL | | **OCR** | DeepSeek-OCR, DeepSeek-OCR-2 , PaddleOCR-VL
PaddleOCR-VL1.5, Hunyuan-OCR, GLM-OCR | | **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR | -| **TTS** | VoxCPM, VoxCPM1.5, VoxCPM2 | +| **TTS** | VoxCPM, VoxCPM1.5, VoxCPM2, Moss-TTS-Nano | | **Image** | RMBG-2.0 (background removal) | | **Embedding** | Qwen3-Embedding, all-MiniLM-L6-v2 | | **Reranker** | Qwen3-Reranker | ## Changelog +### 2026-05-24 +- update doc + ### 2026-05-11 - add Moss-TTS-Nano,its performance is worse than the original Python version @@ -215,27 +218,25 @@ pnpm run tauri build ```rust # VoxCPM example use aha::models::voxcpm::generate::VoxCPMGenerate; -use aha::utils::audio_utils::save_wav; +use aha::utils::audio_utils::save_wav_mono; use anyhow::Result; fn main() -> Result<()> { - let model_path = "xxx/openbmb/VoxCPM-0.5B/"; + let model_path = "xxx/OpenBMB/VoxCPM2/"; let mut voxcpm_generate = VoxCPMGenerate::init(model_path, None, None)?; - - let generate = voxcpm_generate.generate( - "The sun is shining bright, flowers smile at me, birds say early early early".to_string(), + let generate = voxcpm_generate.inference( + "aha是一个基于Rust和Candle框架的本地AI推理引擎,支持多模态模型(文本、视觉、语音、OCR)。".to_string(), None, None, 2, - 100, + 1000, 10, 2.0, - false, 6.0, )?; - let _ = save_wav(&generate, "voxcpm.wav")?; + save_wav_mono(&generate, "voxcpm2.wav", voxcpm_generate.sample_rate() as u32)?; Ok(()) } ``` diff --git a/README.zh-CN.md b/README.zh-CN.md index 4f1879c..5a73dc9 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -38,6 +38,9 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理 | **重排序** | Qwen3-Reranker | ## 更新日志 +### 2026-05-24 +- 更新文档 + ### 2026-05-11 - 添加Moss-TTS-Nano,该模型比较小,计算误差影响较大,效果比python原版差 @@ -214,27 +217,25 @@ pnpm run tauri build ```rust # VoxCPM示例 use aha::models::voxcpm::generate::VoxCPMGenerate; -use aha::utils::audio_utils::save_wav; +use aha::utils::audio_utils::save_wav_mono; use anyhow::Result; fn main() -> Result<()> { - let model_path = "xxx/openbmb/VoxCPM-0.5B/"; - + let model_path = "xxx/OpenBMB/VoxCPM2/"; + let mut voxcpm_generate = VoxCPMGenerate::init(model_path, None, None)?; - - let generate = voxcpm_generate.generate( - "太阳当空照,花儿对我笑,小鸟说早早早".to_string(), + let generate = voxcpm_generate.inference( + "aha是一个基于Rust和Candle框架的本地AI推理引擎,支持多模态模型(文本、视觉、语音、OCR)。".to_string(), None, None, 2, - 100, + 1000, 10, 2.0, - false, 6.0, )?; - let _ = save_wav(&generate, "voxcpm.wav")?; + save_wav_mono(&generate, "voxcpm2.wav", voxcpm_generate.sample_rate() as u32)?; Ok(()) } ``` diff --git a/assets/img/aha_weixinqun.png b/assets/img/aha_weixinqun.png index e997888..be388f0 100644 Binary files a/assets/img/aha_weixinqun.png and b/assets/img/aha_weixinqun.png differ diff --git a/src/models/voxcpm_refact/generate.rs b/src/models/voxcpm_refact/generate.rs index 8162f73..c5bd0f1 100644 --- a/src/models/voxcpm_refact/generate.rs +++ b/src/models/voxcpm_refact/generate.rs @@ -1,11 +1,13 @@ use anyhow::{Result, anyhow}; +use base64::{Engine, prelude::BASE64_STANDARD}; use candle_core::{DType, Device, Tensor, pickle::read_all_with_key}; use candle_nn::VarBuilder; -use rocket::futures::Stream; +use rocket::futures::{Stream, stream}; use std::collections::HashMap; use crate::{ models::{ + GenerateModel, voxcpm::{ audio_vae::AudioVAE, config::{AudioVaeConfig, VoxCPMConfig}, @@ -13,7 +15,12 @@ use crate::{ }, voxcpm_refact::{model::VoxCPMModelRefact, processor::VoxCPMProcessor}, }, - utils::{find_type_files, get_device, get_dtype}, + params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, + utils::{ + audio_utils::{extract_audio_url, get_audio_wav_u8}, + extract_metadata_value, extract_user_text, find_type_files, get_device, get_dtype, + response_utils::build_audio_completion_response, + }, }; pub struct VoxCPMGenerateRefact { @@ -23,7 +30,7 @@ pub struct VoxCPMGenerateRefact { processor: VoxCPMProcessor, prompt_cache: Option>, out_sample_rate: usize, - // model_name: String, + model_name: String, } impl VoxCPMGenerateRefact { @@ -55,11 +62,11 @@ impl VoxCPMGenerateRefact { sr_bin_boundaries: None, }, }; - // let model_name = std::path::Path::new(path) - // .file_name() - // .and_then(|s| s.to_str()) - // .unwrap_or("VoxCPM") - // .to_string(); + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("VoxCPM") + .to_string(); let audio_vae = AudioVAE::new( vb_vae, audio_config.encoder_dim, @@ -118,7 +125,7 @@ impl VoxCPMGenerateRefact { processor, prompt_cache: None, out_sample_rate, - // model_name, + model_name, }) } @@ -126,6 +133,73 @@ impl VoxCPMGenerateRefact { self.out_sample_rate } + pub fn inference( + &mut self, + target_text: String, + prompt_text: Option, + prompt_wav_path: Option, + min_len: usize, + max_len: usize, + inference_timesteps: usize, + cfg_value: f64, + retry_badcase: bool, + retry_badcase_ratio_threshold: f64, + ) -> Result { + let (text_token, audio_feat, audio_mask) = self.processor.processor( + target_text, + prompt_text, + prompt_wav_path, + &self.tokenizer, + &self.audio_vae, + )?; + let target_text_length = if let Some(mask) = &audio_mask { + text_token.dim(1)? - (mask.sum_all()?.to_scalar::()? as usize) + } else { + text_token.dim(1)? + }; + let max_len = if retry_badcase { + (target_text_length as f64 * retry_badcase_ratio_threshold + 10.0) as usize + } else { + max_len + }; + let audio = self.voxcpm.inference( + &text_token, + audio_feat.as_ref(), + audio_mask.as_ref(), + min_len, + max_len, + inference_timesteps, + cfg_value, + &self.audio_vae, + )?; + self.voxcpm.clear_kv_cache(); + Ok(audio) + } + + pub fn generate_with_prompt_simple( + &mut self, + target_text: String, + prompt_text: Option, + prompt_wav_path: Option, + ) -> Result { + let audio = self.inference( + target_text, + prompt_text, + prompt_wav_path, + 2, + 1000, + 10, + 2.0, + false, + 6.0, + )?; + Ok(audio) + } + pub fn generate_simple(&mut self, target_text: String) -> Result { + let audio = self.inference(target_text, None, None, 2, 100, 10, 2.0, false, 6.0)?; + Ok(audio) + } + pub fn build_prompt_cache( &mut self, prompt_text: String, @@ -225,3 +299,79 @@ impl VoxCPMGenerateRefact { } } } + +impl GenerateModel for VoxCPMGenerateRefact { + fn generate(&mut self, mes: ChatCompletionParameters) -> Result { + let prompt_text = extract_metadata_value::(&mes.metadata, "prompt_text"); + let control_instruction = + extract_metadata_value::(&mes.metadata, "control_instruction"); + let min_len = extract_metadata_value::(&mes.metadata, "min_len").unwrap_or(2); + let max_len = extract_metadata_value::(&mes.metadata, "max_len").unwrap_or(4096); + let inference_timesteps = + extract_metadata_value::(&mes.metadata, "inference_timesteps").unwrap_or(10); + let cfg_value = extract_metadata_value::(&mes.metadata, "cfg_value").unwrap_or(2.0); + let retry_badcase_ratio_threshold = + extract_metadata_value::(&mes.metadata, "retry_badcase_ratio_threshold") + .unwrap_or(6.0); + + let prompt_wav = extract_audio_url(&mes); + let prompt_wav_path = if !prompt_wav.is_empty() { + Some(prompt_wav[0].clone()) + } else { + None + }; + if !self.model_name.contains("2") && prompt_wav_path.is_some() && prompt_text.is_none() { + return Err(anyhow!( + "reference mode is only supported with VoxCPM2 models" + )); + } + let mut target_text = extract_user_text(&mes)?; + if let Some(instruction) = control_instruction + && self.model_name.contains("2") + { + target_text = format!("({instruction}){target_text}"); + } + let audio = self + .inference( + target_text, + prompt_text, + prompt_wav_path, + min_len, + max_len, + inference_timesteps, + cfg_value, + true, + retry_badcase_ratio_threshold, + ) + .inspect_err(|_| { + self.voxcpm.clear_kv_cache(); + })?; + let wav_u8 = get_audio_wav_u8(&audio, self.out_sample_rate as u32)?; + // let wave_u8_str = String::from_utf8(wav_u8)?; + let base64_audio = BASE64_STANDARD.encode(wav_u8); + let response = build_audio_completion_response(&base64_audio, &self.model_name); + self.voxcpm.clear_kv_cache(); + Ok(response) + } + #[allow(unused_variables)] + fn generate_stream( + &mut self, + mes: ChatCompletionParameters, + ) -> Result< + Box< + dyn Stream> + + Send + + Unpin + + '_, + >, + > { + let error_stream = stream::once(async { + Err(anyhow::anyhow!(format!( + "{} model not support stream", + self.model_name + ))) as Result + }); + + Ok(Box::new(Box::pin(error_stream))) + } +} diff --git a/src/models/voxcpm_refact/model.rs b/src/models/voxcpm_refact/model.rs index 5ec28de..eb0d7f7 100644 --- a/src/models/voxcpm_refact/model.rs +++ b/src/models/voxcpm_refact/model.rs @@ -196,9 +196,6 @@ impl VoxCPMModelRefact { (text_embed, prefix_feat_cond, None) }; let mut pred_feat_seq = Vec::new(); - // if feat_mask.i((1, t-1))?.to_scalar::()? == 0.0 { - // // TODO for stream - // } let mut position_id = 0; let mut seq_len = t; let enc_outputs = self diff --git a/tests/test_voxcpm.rs b/tests/test_voxcpm.rs index 570e830..d6d7d7f 100644 --- a/tests/test_voxcpm.rs +++ b/tests/test_voxcpm.rs @@ -1,18 +1,16 @@ use std::time::Instant; +use aha::models::voxcpm_refact::generate::VoxCPMGenerateRefact; use aha::params::chat::ChatCompletionParameters; use aha::{ - models::{ - GenerateModel, - voxcpm::{generate::VoxCPMGenerate, tokenizer::SingleChineseTokenizer}, - }, + models::{GenerateModel, voxcpm::tokenizer::SingleChineseTokenizer}, utils::audio_utils::{extract_and_save_audio_from_response, save_wav_mono}, }; use anyhow::{Ok, Result}; #[test] fn voxcpm_use_message_generate() -> Result<()> { - // RUST_BACKTRACE=1 cargo test -F cuda voxcpm_use_message_generate -r -- --nocapture + // RUST_BACKTRACE=1 cargo test -F cuda --test test_voxcpm voxcpm_use_message_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); @@ -42,7 +40,8 @@ fn voxcpm_use_message_generate() -> Result<()> { "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; let i_start = Instant::now(); - let mut voxcpm_generate = VoxCPMGenerate::init(&model_path, None, None)?; + // let mut voxcpm_generate = VoxCPMGenerate::init(&model_path, None, None)?; + let mut voxcpm_generate = VoxCPMGenerateRefact::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); @@ -54,7 +53,6 @@ fn voxcpm_use_message_generate() -> Result<()> { } let i_duration = i_start.elapsed(); println!("Time elapsed in generate is: {:?}", i_duration); - // save_wav_mono(&generate, "voxcpm.wav", 16000)?; Ok(()) } @@ -66,35 +64,16 @@ fn voxcpm_generate() -> Result<()> { let model_path = format!("{}/OpenBMB/VoxCPM-0.5B/", save_dir); let i_start = Instant::now(); - let mut voxcpm_generate = VoxCPMGenerate::init(&model_path, None, None)?; + let mut voxcpm_generate = VoxCPMGenerateRefact::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - // 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 - 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(), + // 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("https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav".to_string()), 2, 100, 10, @@ -103,6 +82,23 @@ fn voxcpm_generate() -> Result<()> { 6.0, )?; + // 创建prompt_cache + // 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_mono(&generate, "voxcpm.wav", 16000)?; @@ -111,6 +107,7 @@ fn voxcpm_generate() -> Result<()> { #[test] fn voxcpm_tokenizer() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda --test test_voxcpm voxcpm_tokenizer -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); diff --git a/tests/test_voxcpm1_5.rs b/tests/test_voxcpm1_5.rs index f3355ba..3ffd396 100644 --- a/tests/test_voxcpm1_5.rs +++ b/tests/test_voxcpm1_5.rs @@ -3,10 +3,7 @@ use std::time::Instant; use aha::models::voxcpm_refact::generate::VoxCPMGenerateRefact; use aha::params::chat::ChatCompletionParameters; use aha::{ - models::{ - GenerateModel, - voxcpm::{generate::VoxCPMGenerate, tokenizer::SingleChineseTokenizer}, - }, + models::{GenerateModel, voxcpm::tokenizer::SingleChineseTokenizer}, utils::audio_utils::{extract_and_save_audio_from_response, save_wav_mono}, }; use anyhow::{Ok, Result}; @@ -43,7 +40,7 @@ fn voxcpm1_5_use_message_generate() -> Result<()> { "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; let i_start = Instant::now(); - let mut voxcpm_generate = VoxCPMGenerate::init(&model_path, None, None)?; + let mut voxcpm_generate = VoxCPMGenerateRefact::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); @@ -66,43 +63,39 @@ fn voxcpm1_5_generate() -> Result<()> { let model_path = format!("{}/OpenBMB/VoxCPM1.5/", save_dir); let i_start = Instant::now(); - let mut voxcpm_generate = VoxCPMGenerate::init(&model_path, None, None)?; + let mut voxcpm_generate = VoxCPMGenerateRefact::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); let i_start = Instant::now(); // let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?; - // let generate = voxcpm_generate.inference( - // "老大爷我来啦,红红火火恍恍惚惚".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, - // 4096, - // 10, - // 2.0, - // // false, - // 6.0, - // )?; - - // 创建prompt_cache - 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( + "老大爷我来啦,红红火火恍恍惚惚".to_string(), + Some("天雷滚滚我好怕怕,劈得我浑身掉渣渣。突破天劫我笑哈哈,逆天改命我吹喇叭,滴答滴答滴滴答".to_string()), + Some("https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav".to_string()), 2, - 100, + 4096, 10, 2.0, false, 6.0, )?; - std::thread::sleep(std::time::Duration::from_secs(2)); + // 创建prompt_cache + // 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, + // 100, + // 10, + // 2.0, + // false, + // 6.0, + // )?; let i_duration = i_start.elapsed(); println!("Time elapsed in generate is: {:?}", i_duration); diff --git a/tests/test_voxcpm2.rs b/tests/test_voxcpm2.rs index f80dfaa..0fd3889 100644 --- a/tests/test_voxcpm2.rs +++ b/tests/test_voxcpm2.rs @@ -1,9 +1,12 @@ use std::time::Instant; use aha::{ - models::{GenerateModel, voxcpm::generate::VoxCPMGenerate}, + models::{ + GenerateModel, voxcpm::generate::VoxCPMGenerate, + voxcpm_refact::generate::VoxCPMGenerateRefact, + }, params::chat::ChatCompletionParameters, - utils::audio_utils::extract_and_save_audio_from_response, + utils::audio_utils::{extract_and_save_audio_from_response, save_wav_mono}, }; use anyhow::Result; @@ -30,26 +33,74 @@ fn voxcpm2_use_message_generate() -> Result<()> { }, { "type": "text", - "text": "你好,这是aha在说话" + "text": "aha是一个基于Rust和Candle框架的本地AI推理引擎,支持多模态模型(文本、视觉、语音、OCR)。" } ] } - ] + ], + "metadata": {"prompt_text": "天雷滚滚我好怕怕,劈得我浑身掉渣渣。突破天劫我笑哈哈,逆天改命我吹喇叭,滴答滴答滴滴答"} } "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut voxcpm_generate = VoxCPMGenerateRefact::init(&model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + // for _ in 0..10 { + // let _ = voxcpm_generate.generate(mes.clone())?; + // } + // let mut times = vec![]; + // for _ in 0..100 { + // let start = Instant::now(); + // let _ = voxcpm_generate.generate(mes.clone())?; + // times.push(start.elapsed()); + // } + // let mean = times.iter().sum::() / 100; + // println!("mean: {:?}", mean); + // times.sort(); + // println!("p99: {:?}", times[99]); + let i_start = Instant::now(); + let generate = voxcpm_generate.generate(mes)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + let save_path = extract_and_save_audio_from_response(&generate, "./")?; + for path in save_path { + println!("save audio: {}", path); + } + Ok(()) +} + +#[test] +fn voxcpm2_generate() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda --test test_voxcpm2 voxcpm2_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/VoxCPM2/", save_dir); + let i_start = Instant::now(); let mut voxcpm_generate = VoxCPMGenerate::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); let i_start = Instant::now(); - let generate = voxcpm_generate.generate(mes)?; - let save_path = extract_and_save_audio_from_response(&generate, "./")?; - for path in save_path { - println!("save audio: {}", path); - } + let generate = voxcpm_generate.inference( + "aha是一个基于Rust和Candle框架的本地AI推理引擎,支持多模态模型(文本、视觉、语音、OCR)。" + .to_string(), + None, + None, + 2, + 1000, + 10, + 2.0, + 6.0, + )?; let i_duration = i_start.elapsed(); println!("Time elapsed in generate is: {:?}", i_duration); + + save_wav_mono( + &generate, + "voxcpm2.wav", + voxcpm_generate.sample_rate() as u32, + )?; Ok(()) }