diff --git a/Cargo.lock b/Cargo.lock index 484227d..44c4e7a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -49,9 +49,8 @@ dependencies = [ [[package]] name = "aha_openai_dive" -version = "1.3.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "45b54fac70f40807364efeb885a33150e163d86b9db687820ba3c5f06aeeaf6e" +version = "1.4.0" +source = "git+https://github.com/jhqxxx/openai-client.git#023a3b07244e73605e89d6d457d76fa24217c9de" dependencies = [ "bytes", "derive_builder", diff --git a/Cargo.toml b/Cargo.toml index 0b77857..d864a72 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -21,7 +21,7 @@ base64 = "0.22.1" num = "0.4.3" minijinja = "2.12.0" tokenizers = "0.22.1" -aha_openai_dive = {version = "1.3.2", features = ["stream"]} +aha_openai_dive = { git ="https://github.com/jhqxxx/openai-client.git", features = ["stream"]} uuid = { version = "1.18.1", features = ["v4"]} chrono = "0.4.42" rocket = { version = "0.5.1", features = ["serde_json", "json"] } diff --git a/README.md b/README.md index e932e55..cd5909e 100644 --- a/README.md +++ b/README.md @@ -111,6 +111,9 @@ cargo run -F cuda -r -- [参数] * deepseek-ocr: deepseek-ai/DeepSeek-OCR 模型 * hunyuan-ocr: Tencent-Hunyuan/HunyuanOCR 模型 * paddleocr-vl: PaddlePaddle/PaddleOCR-VL 模型 + * RMBG2.0: AI-ModelScope/RMBG-2.0 模型 + * voxcpm: OpenBMB/VoxCPM-0.5B 模型 + * voxcpm1.5: OpenBMB/VoxCPM1.5 模型 * 示例:--model deepseek-ocr 或 -m qwen3vl-2b 3. 权重路径 @@ -140,6 +143,34 @@ cargo run -F cuda -r -- [参数] * 如果未指定 --weight-path,程序会自动下载指定模型 * 下载的模型默认保存在 ~/.aha/ 目录下(除非指定了 --save-dir) +#### API接口介绍 +项目提供基于 OpenAI API 兼容的 RESTful 接口,支持多种模型推理任务。 + +##### 接口列表 +1. 对话接口 +- **端点**: `POST /chat/completions` +- **功能**: 多模态对话和文本生成 +- **支持模型**: Qwen2.5VL,Qwen3VL,DeepSeekOCR 等 +- **请求格式**: OpenAI Chat Completion 格式 +- **响应格式**: OpenAI Chat Completion 格式 +- **流式支持**: 支持 + +2. 图像处理接口 +- **端点**: `POST /images/remove_background` +- **功能**: 图像背景移除 +- **支持模型**: RMBG-2.0 +- **请求格式**: OpenAI Chat Completion 格式 +- **响应格式**: OpenAI Chat Completion 格式 +- **流式支持**: 不支持 + +3. 语音生成接口 +- **端点**: `POST /audio/speech` +- **功能**: 语音合成和生成 +- **支持模型**: VoxCPM,VoxCPM1.5 +- **请求格式**: OpenAI Chat Completion 格式 +- **响应格式**: OpenAI Chat Completion 格式 +- **流式支持**: 不支持 + ### 作为库使用 * cargo add aha * 或者在Cargo.toml中添加 @@ -255,6 +286,9 @@ cargo test -F cuda voxcpm_generate -r -- --nocapture 2. 提交新的 Issue,包含详细描述和复现步骤 ## 更新日志 +### v0.1.6 +* 支持RMGB2.0 模型 + ### v0.1.5 * 支持VoxCPM1.5 模型 diff --git a/assets/img/gougou.jpg b/assets/img/gougou.jpg index 827d74a..c4e6252 100644 Binary files a/assets/img/gougou.jpg and b/assets/img/gougou.jpg differ diff --git a/src/api.rs b/src/api.rs index 38e1359..4eb3d58 100644 --- a/src/api.rs +++ b/src/api.rs @@ -104,3 +104,41 @@ pub(crate) async fn chat( } } } + +#[post("/remove_background", data = "")] +pub(crate) async fn remove_background(req: Json) -> (Status, String) { + let response = { + let model_ref = MODEL + .get() + .cloned() + .ok_or_else(|| anyhow::anyhow!("model not init")) + .unwrap(); + model_ref.write().await.generate(req.into_inner()) + }; + match response { + Ok(res) => { + let response_str = serde_json::to_string(&res).unwrap(); + (Status::Ok, response_str) + } + Err(e) => (Status::InternalServerError, e.to_string()), + } +} + +#[post("/speech", data = "")] +pub(crate) async fn speech(req: Json) -> (Status, String) { + let response = { + let model_ref = MODEL + .get() + .cloned() + .ok_or_else(|| anyhow::anyhow!("model not init")) + .unwrap(); + model_ref.write().await.generate(req.into_inner()) + }; + match response { + Ok(res) => { + let response_str = serde_json::to_string(&res).unwrap(); + (Status::Ok, response_str) + } + Err(e) => (Status::InternalServerError, e.to_string()), + } +} diff --git a/src/main.rs b/src/main.rs index 23a0140..9a0241a 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,8 +1,7 @@ use std::time::Duration; -use aha::models::WhichModel; +use aha::{models::WhichModel, utils::get_default_save_dir}; use clap::Parser; -use dirs::home_dir; use modelscope::ModelScope; use rocket::{ Config, @@ -65,13 +64,6 @@ async fn download_model(model_id: &str, save_dir: &str, max_retries: u32) -> any } } -fn get_default_save_dir() -> Option { - home_dir().map(|mut path| { - path.push(".aha"); // 在 home 目录下创建 .aha 文件夹 - path.to_string_lossy().to_string() - }) -} - #[tokio::main] async fn main() -> anyhow::Result<()> { let args = Args::parse(); @@ -86,6 +78,9 @@ async fn main() -> anyhow::Result<()> { WhichModel::DeepSeekOCR => "deepseek-ai/DeepSeek-OCR", WhichModel::HunyuanOCR => "Tencent-Hunyuan/HunyuanOCR", WhichModel::PaddleOCRVL => "PaddlePaddle/PaddleOCR-VL", + WhichModel::RMBG2_0 => "AI-ModelScope/RMBG-2.0", + WhichModel::VoxCPM => "OpenBMB/VoxCPM-0.5B", + WhichModel::VoxCPM1_5 => "OpenBMB/VoxCPM1.5", }; let model_path = match args.weight_path { Some(path) => path, @@ -118,6 +113,10 @@ pub async fn start_http_server(port: u16) -> anyhow::Result<()> { }); builder = builder.mount("/chat", routes![api::chat]); + // /images/remove_background + builder = builder.mount("/images", routes![api::remove_background]); + // /images/speech + builder = builder.mount("/audio", routes![api::speech]); builder.launch().await?; Ok(()) diff --git a/src/models/mod.rs b/src/models/mod.rs index ed5a7f3..705d69e 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -18,7 +18,8 @@ use crate::models::{ deepseek_ocr::generate::DeepseekOCRGenerateModel, hunyuan_ocr::generate::HunyuanOCRGenerateModel, minicpm4::generate::MiniCPMGenerateModel, paddleocr_vl::generate::PaddleOCRVLGenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel, - qwen3vl::generate::Qwen3VLGenerateModel, + qwen3vl::generate::Qwen3VLGenerateModel, rmbg2_0::generate::RMBG2_0Model, + voxcpm::generate::VoxCPMGenerate, }; #[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)] @@ -43,6 +44,12 @@ pub enum WhichModel { HunyuanOCR, #[value(name = "paddleocr-vl")] PaddleOCRVL, + #[value(name = "RMBG2.0")] + RMBG2_0, + #[value(name = "voxcpm")] + VoxCPM, + #[value(name = "voxcpm1.5")] + VoxCPM1_5, } pub trait GenerateModel { @@ -67,6 +74,8 @@ pub enum ModelInstance<'a> { DeepSeekOCR(DeepseekOCRGenerateModel), HunyuanOCR(HunyuanOCRGenerateModel<'a>), PaddleOCRVL(Box>), + RMBG2_0(Box), + VoxCPM(Box), } impl<'a> GenerateModel for ModelInstance<'a> { @@ -78,6 +87,8 @@ impl<'a> GenerateModel for ModelInstance<'a> { ModelInstance::DeepSeekOCR(model) => model.generate(mes), ModelInstance::HunyuanOCR(model) => model.generate(mes), ModelInstance::PaddleOCRVL(model) => model.generate(mes), + ModelInstance::RMBG2_0(model) => model.generate(mes), + ModelInstance::VoxCPM(model) => model.generate(mes), } } @@ -99,6 +110,8 @@ impl<'a> GenerateModel for ModelInstance<'a> { ModelInstance::DeepSeekOCR(model) => model.generate_stream(mes), ModelInstance::HunyuanOCR(model) => model.generate_stream(mes), ModelInstance::PaddleOCRVL(model) => model.generate_stream(mes), + ModelInstance::RMBG2_0(model) => model.generate_stream(mes), + ModelInstance::VoxCPM(model) => model.generate_stream(mes), } } } @@ -145,6 +158,18 @@ pub fn load_model(model_type: WhichModel, path: &str) -> Result { + let model = RMBG2_0Model::init(path, None, None)?; + ModelInstance::RMBG2_0(Box::new(model)) + } + WhichModel::VoxCPM => { + let model = VoxCPMGenerate::init(path, None, None)?; + ModelInstance::VoxCPM(Box::new(model)) + } + WhichModel::VoxCPM1_5 => { + let model = VoxCPMGenerate::init(path, None, None)?; + ModelInstance::VoxCPM(Box::new(model)) + } }; Ok(model) } diff --git a/src/models/rmbg2_0/generate.rs b/src/models/rmbg2_0/generate.rs index 1cff361..57e0a5e 100644 --- a/src/models/rmbg2_0/generate.rs +++ b/src/models/rmbg2_0/generate.rs @@ -1,18 +1,24 @@ -use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use std::io::Cursor; + +use aha_openai_dive::v1::resources::chat::{ + ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, +}; use anyhow::Result; +use base64::{Engine, prelude::BASE64_STANDARD}; use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; use image::{Rgba, RgbaImage}; +use rocket::futures::{Stream, stream}; use crate::{ - models::rmbg2_0::model::BiRefNet, + models::{GenerateModel, rmbg2_0::model::BiRefNet}, utils::{ - find_type_files, get_device, get_dtype, + build_img_completion_response, find_type_files, get_device, get_dtype, img_utils::{extract_images, float_tensor_to_dynamic_image, img_transform_with_resize}, }, }; -pub struct RMBG2_0 { +pub struct RMBG2_0Model { model: BiRefNet, h: u32, w: u32, @@ -20,9 +26,10 @@ pub struct RMBG2_0 { img_std: Tensor, device: Device, dtype: DType, + model_name: String, } -impl RMBG2_0 { +impl RMBG2_0Model { pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { let device = get_device(device); let dtype = get_dtype(dtype, "float32"); @@ -41,10 +48,11 @@ impl RMBG2_0 { img_std, device, dtype, + model_name: "RMBG2.0".to_string(), }) } - pub fn generate(&self, mes: ChatCompletionParameters) -> Result> { + pub fn inference(&self, mes: ChatCompletionParameters) -> Result> { let imgs = extract_images(&mes)?; let mut rmbg_png = vec![]; for img in imgs { @@ -78,3 +86,39 @@ impl RMBG2_0 { Ok(rmbg_png) } } + +impl GenerateModel for RMBG2_0Model { + fn generate(&mut self, mes: ChatCompletionParameters) -> Result { + let rmbg_png = self.inference(mes)?; + let mut base64_vec = vec![]; + for img in rmbg_png { + let mut png_bytes = Vec::new(); + img.write_to(&mut Cursor::new(&mut png_bytes), image::ImageFormat::Png)?; + let base64_string = BASE64_STANDARD.encode(png_bytes); + base64_vec.push(base64_string); + } + let response = build_img_completion_response(&base64_vec, &self.model_name); + 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/generate.rs b/src/models/voxcpm/generate.rs index 4781c5b..97268b2 100644 --- a/src/models/voxcpm/generate.rs +++ b/src/models/voxcpm/generate.rs @@ -1,22 +1,36 @@ use std::collections::HashMap; +use aha_openai_dive::v1::resources::chat::{ + ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, +}; use anyhow::{Ok, Result}; +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, stream}; use crate::{ - models::voxcpm::{ - audio_vae::AudioVAE, - config::{AudioVaeConfig, VoxCPMConfig}, - model::VoxCPMModel, - tokenizer::SingleChineseTokenizer, + models::{ + GenerateModel, + voxcpm::{ + audio_vae::AudioVAE, + config::{AudioVaeConfig, VoxCPMConfig}, + model::VoxCPMModel, + tokenizer::SingleChineseTokenizer, + }, + }, + utils::{ + audio_utils::{extract_audio_url, get_audio_wav_u8}, + build_audio_completion_response, extract_metadata_value, extract_user_text, + find_type_files, get_device, get_dtype, }, - utils::{find_type_files, get_device, get_dtype}, }; pub struct VoxCPMGenerate { voxcpm: VoxCPMModel, prompt_cache: Option>, + sample_rate: usize, + model_name: String, } impl VoxCPMGenerate { @@ -48,6 +62,11 @@ impl VoxCPMGenerate { sample_rate: 16000, }, }; + let model_name = if audio_config.sample_rate == 16000 { + "VoxCPM".to_string() + } else { + "VoxCPM1.5".to_string() + }; let audio_vae = AudioVAE::new( vb_vae, audio_config.encoder_dim, @@ -85,6 +104,8 @@ impl VoxCPMGenerate { Ok(Self { voxcpm, prompt_cache: None, + sample_rate: audio_config.sample_rate, + model_name, }) } @@ -135,7 +156,7 @@ impl VoxCPMGenerate { prompt_text: Option, prompt_wav_path: Option, ) -> Result { - let audio = self.generate( + let audio = self.inference( target_text, prompt_text, prompt_wav_path, @@ -150,10 +171,10 @@ impl VoxCPMGenerate { } pub fn generate_simple(&mut self, target_text: String) -> Result { // let audio = self.generate(target_text, None, None, 2, 100, 10, 2.0, false, 6.0)?; - let audio = self.generate(target_text, None, None, 2, 100, 10, 2.0, 6.0)?; + let audio = self.inference(target_text, None, None, 2, 100, 10, 2.0, 6.0)?; Ok(audio) } - pub fn generate( + pub fn inference( &mut self, target_text: String, prompt_text: Option, @@ -179,3 +200,59 @@ impl VoxCPMGenerate { Ok(audio) } } + +impl GenerateModel for VoxCPMGenerate { + fn generate(&mut self, mes: ChatCompletionParameters) -> Result { + let prompt_text = extract_metadata_value::(&mes.metadata, "prompt_text"); + 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 target_text = extract_user_text(&mes)?; + let prompt_wav = extract_audio_url(&mes)?; + let prompt_wav_path = if !prompt_wav.is_empty() { + Some(prompt_wav[0].clone()) + } else { + None + }; + let audio = self.voxcpm.generate( + target_text, + prompt_text, + prompt_wav_path, + min_len, + max_len, + inference_timesteps, + cfg_value, + retry_badcase_ratio_threshold, + )?; + let wav_u8 = get_audio_wav_u8(&audio, self.sample_rate as u32)?; + let base64_audio = BASE64_STANDARD.encode(wav_u8); + let response = build_audio_completion_response(&base64_audio, &self.model_name); + 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/model.rs b/src/models/voxcpm/model.rs index cd481c8..08135c9 100644 --- a/src/models/voxcpm/model.rs +++ b/src/models/voxcpm/model.rs @@ -513,7 +513,7 @@ impl VoxCPMModel { let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?; let text_length = text_token.dim(0)?; let mut audio = - load_audio_with_resample(path, self.device.clone(), Some(self.sample_rate))?; + load_audio_with_resample(&path, self.device.clone(), Some(self.sample_rate))?; let patch_len = self.patch_size * self.chunk_size; if audio.dim(1)? % patch_len != 0 { audio = audio.pad_with_zeros( @@ -728,8 +728,11 @@ impl VoxCPMModel { ) -> Result> { let text_token = self.tokenizer.encode(prompt_text)?; let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?; - let mut audio = - load_audio_with_resample(prompt_wav_path, self.device.clone(), Some(self.sample_rate))?; + let mut audio = load_audio_with_resample( + &prompt_wav_path, + self.device.clone(), + Some(self.sample_rate), + )?; let patch_len = self.patch_size * self.chunk_size; if audio.dim(1)? % patch_len != 0 { audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?; diff --git a/src/utils/audio_utils.rs b/src/utils/audio_utils.rs index a0ac671..33af770 100644 --- a/src/utils/audio_utils.rs +++ b/src/utils/audio_utils.rs @@ -1,12 +1,22 @@ -use std::f64::consts::PI; -use std::path::Path; +use std::fs::File; +use std::io::Write; +use std::path::{Path, PathBuf}; +use std::{f64::consts::PI, io::Cursor}; +use aha_openai_dive::v1::resources::chat::{ + ChatCompletionParameters, ChatCompletionResponse, ChatMessage, ChatMessageContent, + ChatMessageContentPart, +}; use anyhow::{Result, anyhow}; +use base64::Engine; +use base64::prelude::BASE64_STANDARD; use candle_core::{D, Device, Tensor}; use candle_nn::{Conv1d, Conv1dConfig, Module}; use hound::{SampleFormat, WavReader}; use num::integer::gcd; +use crate::utils::get_default_save_dir; + // 重采样方法枚举 #[derive(Debug, Clone, Copy)] pub enum ResamplingMethod { @@ -213,9 +223,62 @@ pub fn resample_simple(waveform: &Tensor, orig_freq: i64, new_freq: i64) -> Resu None, ) } +pub fn load_audio_from_url(url: &str) -> Result { + tokio::task::block_in_place(|| { + let client = reqwest::blocking::Client::new(); + let response = client.get(url).send()?; + if !response.status().is_success() { + return Err(anyhow::anyhow!( + "Failed to download file: {}", + response.status() + )); + } + let temp_dir = get_default_save_dir().expect("Failed to get home directory"); + let temp_dir = PathBuf::from(temp_dir); + let temp_path = temp_dir.join("temp_audio.wav"); -pub fn load_audio>(path: P, device: Device) -> Result<(Tensor, usize)> { - let mut reader = WavReader::open(path)?; + let mut file = std::fs::File::create(&temp_path)?; + let mut content = Cursor::new(response.bytes()?); + std::io::copy(&mut content, &mut file)?; + + // Return the temp directory to keep it alive until the function ends + Ok(temp_path) + }) +} + +pub fn get_audio_path(path_str: &str) -> Result { + if path_str.starts_with("http://") || path_str.starts_with("https://") { + // Download file from network + load_audio_from_url(path_str) + } else if path_str.starts_with("file://") { + // Convert file:// URL to local path + let path = url::Url::parse(path_str)?; + let path = path.to_file_path(); + let path = match path { + Ok(path) => path, + Err(_) => { + let mut path = path_str.to_owned(); + path = path.split_off(7); + PathBuf::from(path) + } + }; + Ok(path) + } else if path_str.starts_with("data:audio") && path_str.contains("base64,") { + let data: Vec<&str> = path_str.split("base64,").collect(); + let data = data[1]; + let temp_dir = get_default_save_dir().expect("Failed to get home directory"); + let temp_dir = PathBuf::from(temp_dir); + let temp_path = temp_dir.join("temp_audio.wav"); + save_audio_from_base64(data, &temp_path)?; + Ok(temp_path) + } else { + Err(anyhow::anyhow!("get audio path error {}", path_str)) + } +} + +pub fn load_audio(path: &str, device: Device) -> Result<(Tensor, usize)> { + let audio_path = get_audio_path(path)?; + let mut reader = WavReader::open(audio_path)?; let spec = reader.spec(); let samples: Vec = match spec.sample_format { SampleFormat::Int => { @@ -264,8 +327,8 @@ pub fn load_audio>(path: P, device: Device) -> Result<(Tensor, us Ok((audio_tensor, sample_rate as usize)) } -pub fn load_audio_with_resample>( - path: P, +pub fn load_audio_with_resample( + path: &str, device: Device, target_sample_rate: Option, ) -> Result { @@ -299,3 +362,114 @@ pub fn save_wav(audio: &Tensor, save_path: &str, sample_rate: u32) -> Result<()> writer.finalize().unwrap(); Ok(()) } + +pub fn get_audio_wav_u8(audio: &Tensor, sample_rate: u32) -> Result> { + let spec = hound::WavSpec { + channels: 1, + sample_rate, + bits_per_sample: 16, + sample_format: hound::SampleFormat::Int, + }; + assert_eq!(audio.dim(0)?, 1, "audio channel must be 1"); + let max = audio.abs()?.max_all()?; + let max = max.to_scalar::()?; + let ratio = if max > 1.0 { 32767.0 / max } else { 32767.0 }; + let audio = audio.squeeze(0)?; + let audio_vec = audio.to_vec1::()?; + let mut cursor = Cursor::new(Vec::new()); + let mut writer = hound::WavWriter::new(&mut cursor, spec)?; + for i in audio_vec { + let sample_i16 = (i * ratio).round() as i16; + writer.write_sample(sample_i16)?; + } + writer.finalize()?; + let wav_buffer = cursor.into_inner(); + Ok(wav_buffer) +} + +pub fn extract_audio_url(mes: &ChatCompletionParameters) -> Result> { + let mut audio_vec = Vec::new(); + for chat_mes in mes.messages.clone() { + if let ChatMessage::User { content, .. } = chat_mes.clone() + && let ChatMessageContent::ContentPart(part_vec) = content + { + for part in part_vec { + if let ChatMessageContentPart::Audio(audio_part) = part { + let audio_url = audio_part.audio_url; + audio_vec.push(audio_url.url); + } + } + } + // if let ChatMessage::User { content, .. } = chat_mes.clone() + // && let ChatMessageContent::ContentPart(part_vec) = content + // { + // for part in part_vec { + // if let ChatMessageContentPart::Text(text_part) = part { + // let text = text_part.text; + // if text.chars().count() > 0 { + // ret = ret + &text + "\n" + // } + // } + // } + // } + } + Ok(audio_vec) +} + +// 从 ChatCompletionResponse 中提取音频数据 +pub fn extract_audio_base64_from_response( + response: &ChatCompletionResponse, +) -> Result> { + let mut audio_data_list = Vec::new(); + + for choice in &response.choices { + if let ChatMessage::Assistant { + content: Some(ChatMessageContent::ContentPart(parts)), + .. + } = &choice.message + { + for part in parts.clone() { + if let ChatMessageContentPart::Audio(audio_part) = part { + // if let Some(audio_data) = &audio_part.audio_url { + // audio_data_list.push(audio_data.data.clone()); + // } + let audio_url = audio_part.audio_url; + audio_data_list.push(audio_url.url); + } + } + } + } + + Ok(audio_data_list) +} + +// 将 base64 音频数据解码并保存到文件 +pub fn save_audio_from_base64>(base64_data: &str, file_path: P) -> Result<()> { + // 解码 base64 数据 + let data: Vec<&str> = base64_data.split("base64,").collect(); + let data = data[1]; + let decoded_data = BASE64_STANDARD.decode(data)?; + + // 创建文件并写入数据 + let mut file = File::create(file_path)?; + file.write_all(&decoded_data)?; + + Ok(()) +} + +// 组合函数:从响应中提取音频并保存到文件 +pub fn extract_and_save_audio_from_response( + response: &ChatCompletionResponse, + directory: &str, +) -> Result> { + let audio_data_list = extract_audio_base64_from_response(response)?; + let mut saved_files = Vec::new(); + + for (index, audio_data) in audio_data_list.iter().enumerate() { + let file_path = format!("{}/audio_{}.wav", directory, index); + save_audio_from_base64(audio_data, &file_path)?; + saved_files.push(file_path); + } + + Ok(saved_files) +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 95a45c2..d14996e 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -3,19 +3,21 @@ pub mod img_utils; pub mod tensor_utils; pub mod video_utils; -use std::process::Command; +use std::{fs, process::Command}; use aha_openai_dive::v1::resources::{ chat::{ - ChatCompletionChoice, ChatCompletionChunkChoice, ChatCompletionChunkResponse, - ChatCompletionParameters, ChatCompletionResponse, ChatMessage, ChatMessageContent, - ChatMessageContentPart, DeltaChatMessage, DeltaFunction, DeltaToolCall, Function, ToolCall, + AudioUrlType, ChatCompletionChoice, ChatCompletionChunkChoice, ChatCompletionChunkResponse, + ChatCompletionParameters, ChatCompletionResponse, ChatMessage, ChatMessageAudioContentPart, + ChatMessageContent, ChatMessageContentPart, ChatMessageImageContentPart, DeltaChatMessage, + DeltaFunction, DeltaToolCall, Function, ImageUrlType, ToolCall, }, shared::{FinishReason, Usage}, }; use anyhow::Result; use candle_core::{DType, Device}; use candle_transformers::generation::{LogitsProcessor, Sampling}; +use dirs::home_dir; pub fn get_device(device: Option<&Device>) -> Device { match device { @@ -137,6 +139,90 @@ pub fn ceil_by_factor(num: f32, factor: u32) -> u32 { ceil * factor } +pub fn build_img_completion_response( + base64vec: &Vec, + model_name: &str, +) -> ChatCompletionResponse { + let id = uuid::Uuid::new_v4().to_string(); + let mut response = ChatCompletionResponse { + id: Some(id), + choices: vec![], + created: chrono::Utc::now().timestamp() as u32, + model: model_name.to_string(), + service_tier: None, + system_fingerprint: None, + object: "chat.completion".to_string(), + usage: None, + }; + let mut conten_part_vec = vec![]; + for img_bas64 in base64vec { + let img_base64_prefix = "data:image/png;base64,".to_string() + img_bas64; + let part = ChatMessageContentPart::Image(ChatMessageImageContentPart { + r#type: "image".to_string(), + image_url: ImageUrlType { + url: img_base64_prefix, + detail: None, + }, + }); + conten_part_vec.push(part); + } + let choice = ChatCompletionChoice { + index: 0, + message: ChatMessage::Assistant { + content: Some(ChatMessageContent::ContentPart(conten_part_vec)), + reasoning_content: None, + refusal: None, + name: None, + audio: None, + tool_calls: None, + }, + finish_reason: Some(FinishReason::StopSequenceReached), + logprobs: None, + }; + response.choices.push(choice); + response +} + +pub fn build_audio_completion_response( + base64_audio: &String, + model_name: &str, +) -> ChatCompletionResponse { + let id = uuid::Uuid::new_v4().to_string(); + let mut response = ChatCompletionResponse { + id: Some(id), + choices: vec![], + created: chrono::Utc::now().timestamp() as u32, + model: model_name.to_string(), + service_tier: None, + system_fingerprint: None, + object: "chat.completion".to_string(), + usage: None, + }; + + let base64_audio = format!("data:audio/wav;base64,{}", base64_audio); + let conten_part_vec = vec![ChatMessageContentPart::Audio(ChatMessageAudioContentPart { + r#type: "audio".to_string(), + audio_url: AudioUrlType { + url: base64_audio.to_string(), + }, + })]; + let choice = ChatCompletionChoice { + index: 0, + message: ChatMessage::Assistant { + content: Some(ChatMessageContent::ContentPart(conten_part_vec)), + reasoning_content: None, + refusal: None, + name: None, + audio: None, + tool_calls: None, + }, + finish_reason: Some(FinishReason::StopSequenceReached), + logprobs: None, + }; + response.choices.push(choice); + response +} + pub fn build_completion_response( res: String, model_name: &str, @@ -144,10 +230,6 @@ pub fn build_completion_response( ) -> ChatCompletionResponse { let id = uuid::Uuid::new_v4().to_string(); let usage = num_tokens.map(|num| Usage { - input_tokens: None, - input_tokens_details: None, - output_tokens: None, - output_tokens_details: None, prompt_tokens: None, completion_tokens: None, total_tokens: num, @@ -358,3 +440,49 @@ pub fn extract_mes(mes: &ChatCompletionParameters) -> Result( + metadata: &Option>, + key: &str, +) -> Option +where + T: std::str::FromStr + Clone + PartialEq, +{ + if let Some(map) = metadata + && let Some(value_str) = map.get(key) + && let Ok(value) = value_str.parse::() + { + return Some(value); + } + None +} + +pub fn extract_user_text(mes: &ChatCompletionParameters) -> Result { + let mut ret = "".to_string(); + for chat_mes in mes.messages.clone() { + if let ChatMessage::User { content, .. } = chat_mes.clone() + && let ChatMessageContent::ContentPart(part_vec) = content + { + for part in part_vec { + if let ChatMessageContentPart::Text(text_part) = part { + let text = text_part.text; + if text.chars().count() > 0 { + ret = ret + &text + "\n" + } + } + } + } + } + ret = ret.trim().to_string(); + Ok(ret) +} + +pub fn get_default_save_dir() -> Option { + home_dir().map(|mut path| { + path.push(".aha"); + if let Err(e) = fs::create_dir_all(&path) { + eprintln!("Failed to create directory {:?}: {}", path, e); + } + path.to_string_lossy().to_string() + }) +} diff --git a/tests/messy_test.rs b/tests/messy_test.rs index b0006f5..9a0a5cc 100644 --- a/tests/messy_test.rs +++ b/tests/messy_test.rs @@ -1,3 +1,4 @@ +use aha::utils::get_default_save_dir; use anyhow::Result; use candle_core::Tensor; @@ -5,18 +6,19 @@ use candle_core::Tensor; fn messy_test() -> Result<()> { // RUST_BACKTRACE=1 cargo test -F cuda messy_test -r -- --nocapture let device = &candle_core::Device::Cpu; + // let path = get_default_save_dir(); let x = Tensor::arange(0.0, 9.0, device)?; println!("x: {}", x); - let x = x - .unsqueeze(0)? - .unsqueeze(0)? - .broadcast_as((5, 5, 9))? - .reshape((5, 5, 3, 3))?; - println!("x: {}", x); - let x = x.permute((0, 2, 1, 3))?; - println!("x: {}", x); - let x = x.reshape((15, 15))?; - println!("x: {}", x); + // let x = x + // .unsqueeze(0)? + // .unsqueeze(0)? + // .broadcast_as((5, 5, 9))? + // .reshape((5, 5, 3, 3))?; + // println!("x: {}", x); + // let x = x.permute((0, 2, 1, 3))?; + // println!("x: {}", x); + // let x = x.reshape((15, 15))?; + // println!("x: {}", x); // let xs = Tensor::rand(0.0, 5.0, (1, 1, 3, 3), device)?; // println!("xs: {}", xs); // let xs = xs.pad_with_zeros(3, 2, 2)? diff --git a/tests/test_rmbg2_0.rs b/tests/test_rmbg2_0.rs index 6a4364a..ef75833 100644 --- a/tests/test_rmbg2_0.rs +++ b/tests/test_rmbg2_0.rs @@ -1,6 +1,6 @@ use std::time::Instant; -use aha::models::rmbg2_0::generate::RMBG2_0; +use aha::models::rmbg2_0::generate::RMBG2_0Model; use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; use anyhow::Result; @@ -31,12 +31,12 @@ fn rmbg2_0_generate() -> Result<()> { "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; let i_start = Instant::now(); - let model = RMBG2_0::init(model_path, None, None)?; + let model = RMBG2_0Model::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 result = model.generate(mes)?; + let result = model.inference(mes)?; let i_duration = i_start.elapsed(); println!("Time elapsed in generate is: {:?}", i_duration); for (i, img) in result.iter().enumerate() { diff --git a/tests/test_voxcpm.rs b/tests/test_voxcpm.rs index dddd2d1..783fb7c 100644 --- a/tests/test_voxcpm.rs +++ b/tests/test_voxcpm.rs @@ -1,14 +1,64 @@ use std::time::Instant; use aha::{ - models::voxcpm::{generate::VoxCPMGenerate, tokenizer::SingleChineseTokenizer}, - utils::audio_utils::save_wav, + models::{ + GenerateModel, + voxcpm::{generate::VoxCPMGenerate, tokenizer::SingleChineseTokenizer}, + }, + utils::audio_utils::{extract_and_save_audio_from_response, save_wav}, }; +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; use anyhow::{Ok, Result}; +#[test] +fn voxcpm_use_message_generate() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda voxcpm_use_message_generate -r -- --nocapture + let model_path = "/home/jhq/huggingface_model/openbmb/VoxCPM-0.5B/"; + let message = r#" + { + "model": "voxcpm", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "audio", + "audio_url": + { + "url": "https://sis-sample-audio.obs.cn-north-1.myhuaweicloud.com/16k16bit.wav" + } + }, + { + "type": "text", + "text": "VoxCPM is an innovative end-to-end TTS model from ModelBest, designed to generate highly realistic speech." + } + ] + } + ], + "metadata": {"prompt_text": "华为致力于把数字世界带给每个人,每个家庭,每个组织,构建万物互联的智能世界。"} + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + 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 i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + // save_wav(&generate, "voxcpm.wav", 16000)?; + Ok(()) +} + #[test] fn voxcpm_generate() -> Result<()> { - // RUST_BACKTRACE=1 cargo test -F cuda,flash-attn voxcpm_generate -r -- --nocapture + // RUST_BACKTRACE=1 cargo test -F cuda voxcpm_generate -r -- --nocapture let model_path = "/home/jhq/huggingface_model/openbmb/VoxCPM-0.5B/"; let i_start = Instant::now(); @@ -18,12 +68,12 @@ fn voxcpm_generate() -> Result<()> { let i_start = Instant::now(); // let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?; - let generate = voxcpm_generate.generate( + 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("./assets/audio/voice_01.wav".to_string()), + Some("file://./assets/audio/voice_01.wav".to_string()), // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()), - // Some("./assets/audio/voice_05.wav".to_string()), + // Some("file://./assets/audio/voice_05.wav".to_string()), 2, 100, 10, @@ -35,7 +85,7 @@ fn voxcpm_generate() -> Result<()> { // 创建prompt_cache // let _ = voxcpm_generate.build_prompt_cache( // "啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(), - // "./assets/audio/voice_01.wav".to_string(), + // "file://./assets/audio/voice_01.wav".to_string(), // )?; // // 使用prompt_cache生成语音 // let generate = voxcpm_generate.generate_use_prompt_cache( diff --git a/tests/test_voxcpm1_5.rs b/tests/test_voxcpm1_5.rs index cf2dca7..da1111a 100644 --- a/tests/test_voxcpm1_5.rs +++ b/tests/test_voxcpm1_5.rs @@ -1,11 +1,60 @@ use std::time::Instant; use aha::{ - models::voxcpm::{generate::VoxCPMGenerate, tokenizer::SingleChineseTokenizer}, - utils::audio_utils::save_wav, + models::{ + GenerateModel, + voxcpm::{generate::VoxCPMGenerate, tokenizer::SingleChineseTokenizer}, + }, + utils::audio_utils::{extract_and_save_audio_from_response, save_wav}, }; +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; use anyhow::{Ok, Result}; +#[test] +fn voxcpm1_5_use_message_generate() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda voxcpm1_5_use_message_generate -r -- --nocapture + let model_path = "/home/jhq/huggingface_model/OpenBMB/VoxCPM1.5/"; + let message = r#" + { + "model": "voxcpm1.5", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "audio", + "audio_url": + { + "url": "https://sis-sample-audio.obs.cn-north-1.myhuaweicloud.com/16k16bit.wav" + } + }, + { + "type": "text", + "text": "VoxCPM is an innovative end-to-end TTS model from ModelBest, designed to generate highly realistic speech." + } + ] + } + ], + "metadata": {"prompt_text": "华为致力于把数字世界带给每个人,每个家庭,每个组织,构建万物互联的智能世界。"} + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + 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 i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + Ok(()) +} + #[test] fn voxcpm1_5_generate() -> Result<()> { // RUST_BACKTRACE=1 cargo test -F cuda voxcpm1_5_generate -r -- --nocapture @@ -18,12 +67,12 @@ fn voxcpm1_5_generate() -> Result<()> { let i_start = Instant::now(); // let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?; - let generate = voxcpm_generate.generate( + 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("./assets/audio/voice_01.wav".to_string()), + Some("file://./assets/audio/voice_01.wav".to_string()), // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()), - // Some("./assets/audio/voice_05.wav".to_string()), + // Some("file://./assets/audio/voice_05.wav".to_string()), 2, 4096, 10, @@ -35,7 +84,7 @@ fn voxcpm1_5_generate() -> Result<()> { // 创建prompt_cache // let _ = voxcpm_generate.build_prompt_cache( // "啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(), - // "./assets/audio/voice_01.wav".to_string(), + // "file://./assets/audio/voice_01.wav".to_string(), // )?; // // 使用prompt_cache生成语音 // let generate = voxcpm_generate.generate_use_prompt_cache(