add image api and audio api
This commit is contained in:
Generated
+2
-3
@@ -49,9 +49,8 @@ dependencies = [
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "aha_openai_dive"
|
name = "aha_openai_dive"
|
||||||
version = "1.3.2"
|
version = "1.4.0"
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "git+https://github.com/jhqxxx/openai-client.git#023a3b07244e73605e89d6d457d76fa24217c9de"
|
||||||
checksum = "45b54fac70f40807364efeb885a33150e163d86b9db687820ba3c5f06aeeaf6e"
|
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"bytes",
|
"bytes",
|
||||||
"derive_builder",
|
"derive_builder",
|
||||||
|
|||||||
+1
-1
@@ -21,7 +21,7 @@ base64 = "0.22.1"
|
|||||||
num = "0.4.3"
|
num = "0.4.3"
|
||||||
minijinja = "2.12.0"
|
minijinja = "2.12.0"
|
||||||
tokenizers = "0.22.1"
|
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"]}
|
uuid = { version = "1.18.1", features = ["v4"]}
|
||||||
chrono = "0.4.42"
|
chrono = "0.4.42"
|
||||||
rocket = { version = "0.5.1", features = ["serde_json", "json"] }
|
rocket = { version = "0.5.1", features = ["serde_json", "json"] }
|
||||||
|
|||||||
@@ -111,6 +111,9 @@ cargo run -F cuda -r -- [参数]
|
|||||||
* deepseek-ocr: deepseek-ai/DeepSeek-OCR 模型
|
* deepseek-ocr: deepseek-ai/DeepSeek-OCR 模型
|
||||||
* hunyuan-ocr: Tencent-Hunyuan/HunyuanOCR 模型
|
* hunyuan-ocr: Tencent-Hunyuan/HunyuanOCR 模型
|
||||||
* paddleocr-vl: PaddlePaddle/PaddleOCR-VL 模型
|
* 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
|
* 示例:--model deepseek-ocr 或 -m qwen3vl-2b
|
||||||
|
|
||||||
3. 权重路径
|
3. 权重路径
|
||||||
@@ -140,6 +143,34 @@ cargo run -F cuda -r -- [参数]
|
|||||||
* 如果未指定 --weight-path,程序会自动下载指定模型
|
* 如果未指定 --weight-path,程序会自动下载指定模型
|
||||||
* 下载的模型默认保存在 ~/.aha/ 目录下(除非指定了 --save-dir)
|
* 下载的模型默认保存在 ~/.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 add aha
|
||||||
* 或者在Cargo.toml中添加
|
* 或者在Cargo.toml中添加
|
||||||
@@ -255,6 +286,9 @@ cargo test -F cuda voxcpm_generate -r -- --nocapture
|
|||||||
2. 提交新的 Issue,包含详细描述和复现步骤
|
2. 提交新的 Issue,包含详细描述和复现步骤
|
||||||
|
|
||||||
## 更新日志
|
## 更新日志
|
||||||
|
### v0.1.6
|
||||||
|
* 支持RMGB2.0 模型
|
||||||
|
|
||||||
### v0.1.5
|
### v0.1.5
|
||||||
* 支持VoxCPM1.5 模型
|
* 支持VoxCPM1.5 模型
|
||||||
|
|
||||||
|
|||||||
Binary file not shown.
|
Before Width: | Height: | Size: 18 MiB After Width: | Height: | Size: 149 KiB |
+38
@@ -104,3 +104,41 @@ pub(crate) async fn chat(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[post("/remove_background", data = "<req>")]
|
||||||
|
pub(crate) async fn remove_background(req: Json<ChatCompletionParameters>) -> (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 = "<req>")]
|
||||||
|
pub(crate) async fn speech(req: Json<ChatCompletionParameters>) -> (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()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+8
-9
@@ -1,8 +1,7 @@
|
|||||||
use std::time::Duration;
|
use std::time::Duration;
|
||||||
|
|
||||||
use aha::models::WhichModel;
|
use aha::{models::WhichModel, utils::get_default_save_dir};
|
||||||
use clap::Parser;
|
use clap::Parser;
|
||||||
use dirs::home_dir;
|
|
||||||
use modelscope::ModelScope;
|
use modelscope::ModelScope;
|
||||||
use rocket::{
|
use rocket::{
|
||||||
Config,
|
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<String> {
|
|
||||||
home_dir().map(|mut path| {
|
|
||||||
path.push(".aha"); // 在 home 目录下创建 .aha 文件夹
|
|
||||||
path.to_string_lossy().to_string()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::main]
|
#[tokio::main]
|
||||||
async fn main() -> anyhow::Result<()> {
|
async fn main() -> anyhow::Result<()> {
|
||||||
let args = Args::parse();
|
let args = Args::parse();
|
||||||
@@ -86,6 +78,9 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
WhichModel::DeepSeekOCR => "deepseek-ai/DeepSeek-OCR",
|
WhichModel::DeepSeekOCR => "deepseek-ai/DeepSeek-OCR",
|
||||||
WhichModel::HunyuanOCR => "Tencent-Hunyuan/HunyuanOCR",
|
WhichModel::HunyuanOCR => "Tencent-Hunyuan/HunyuanOCR",
|
||||||
WhichModel::PaddleOCRVL => "PaddlePaddle/PaddleOCR-VL",
|
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 {
|
let model_path = match args.weight_path {
|
||||||
Some(path) => 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]);
|
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?;
|
builder.launch().await?;
|
||||||
Ok(())
|
Ok(())
|
||||||
|
|||||||
+26
-1
@@ -18,7 +18,8 @@ use crate::models::{
|
|||||||
deepseek_ocr::generate::DeepseekOCRGenerateModel,
|
deepseek_ocr::generate::DeepseekOCRGenerateModel,
|
||||||
hunyuan_ocr::generate::HunyuanOCRGenerateModel, minicpm4::generate::MiniCPMGenerateModel,
|
hunyuan_ocr::generate::HunyuanOCRGenerateModel, minicpm4::generate::MiniCPMGenerateModel,
|
||||||
paddleocr_vl::generate::PaddleOCRVLGenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel,
|
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)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)]
|
||||||
@@ -43,6 +44,12 @@ pub enum WhichModel {
|
|||||||
HunyuanOCR,
|
HunyuanOCR,
|
||||||
#[value(name = "paddleocr-vl")]
|
#[value(name = "paddleocr-vl")]
|
||||||
PaddleOCRVL,
|
PaddleOCRVL,
|
||||||
|
#[value(name = "RMBG2.0")]
|
||||||
|
RMBG2_0,
|
||||||
|
#[value(name = "voxcpm")]
|
||||||
|
VoxCPM,
|
||||||
|
#[value(name = "voxcpm1.5")]
|
||||||
|
VoxCPM1_5,
|
||||||
}
|
}
|
||||||
|
|
||||||
pub trait GenerateModel {
|
pub trait GenerateModel {
|
||||||
@@ -67,6 +74,8 @@ pub enum ModelInstance<'a> {
|
|||||||
DeepSeekOCR(DeepseekOCRGenerateModel),
|
DeepSeekOCR(DeepseekOCRGenerateModel),
|
||||||
HunyuanOCR(HunyuanOCRGenerateModel<'a>),
|
HunyuanOCR(HunyuanOCRGenerateModel<'a>),
|
||||||
PaddleOCRVL(Box<PaddleOCRVLGenerateModel<'a>>),
|
PaddleOCRVL(Box<PaddleOCRVLGenerateModel<'a>>),
|
||||||
|
RMBG2_0(Box<RMBG2_0Model>),
|
||||||
|
VoxCPM(Box<VoxCPMGenerate>),
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<'a> GenerateModel for ModelInstance<'a> {
|
impl<'a> GenerateModel for ModelInstance<'a> {
|
||||||
@@ -78,6 +87,8 @@ impl<'a> GenerateModel for ModelInstance<'a> {
|
|||||||
ModelInstance::DeepSeekOCR(model) => model.generate(mes),
|
ModelInstance::DeepSeekOCR(model) => model.generate(mes),
|
||||||
ModelInstance::HunyuanOCR(model) => model.generate(mes),
|
ModelInstance::HunyuanOCR(model) => model.generate(mes),
|
||||||
ModelInstance::PaddleOCRVL(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::DeepSeekOCR(model) => model.generate_stream(mes),
|
||||||
ModelInstance::HunyuanOCR(model) => model.generate_stream(mes),
|
ModelInstance::HunyuanOCR(model) => model.generate_stream(mes),
|
||||||
ModelInstance::PaddleOCRVL(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<ModelInstance<'_
|
|||||||
let model = PaddleOCRVLGenerateModel::init(path, None, None)?;
|
let model = PaddleOCRVLGenerateModel::init(path, None, None)?;
|
||||||
ModelInstance::PaddleOCRVL(Box::new(model))
|
ModelInstance::PaddleOCRVL(Box::new(model))
|
||||||
}
|
}
|
||||||
|
WhichModel::RMBG2_0 => {
|
||||||
|
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)
|
Ok(model)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 anyhow::Result;
|
||||||
|
use base64::{Engine, prelude::BASE64_STANDARD};
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device, Tensor};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use image::{Rgba, RgbaImage};
|
use image::{Rgba, RgbaImage};
|
||||||
|
use rocket::futures::{Stream, stream};
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::rmbg2_0::model::BiRefNet,
|
models::{GenerateModel, rmbg2_0::model::BiRefNet},
|
||||||
utils::{
|
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},
|
img_utils::{extract_images, float_tensor_to_dynamic_image, img_transform_with_resize},
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct RMBG2_0 {
|
pub struct RMBG2_0Model {
|
||||||
model: BiRefNet,
|
model: BiRefNet,
|
||||||
h: u32,
|
h: u32,
|
||||||
w: u32,
|
w: u32,
|
||||||
@@ -20,9 +26,10 @@ pub struct RMBG2_0 {
|
|||||||
img_std: Tensor,
|
img_std: Tensor,
|
||||||
device: Device,
|
device: Device,
|
||||||
dtype: DType,
|
dtype: DType,
|
||||||
|
model_name: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl RMBG2_0 {
|
impl RMBG2_0Model {
|
||||||
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
|
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
|
||||||
let device = get_device(device);
|
let device = get_device(device);
|
||||||
let dtype = get_dtype(dtype, "float32");
|
let dtype = get_dtype(dtype, "float32");
|
||||||
@@ -41,10 +48,11 @@ impl RMBG2_0 {
|
|||||||
img_std,
|
img_std,
|
||||||
device,
|
device,
|
||||||
dtype,
|
dtype,
|
||||||
|
model_name: "RMBG2.0".to_string(),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn generate(&self, mes: ChatCompletionParameters) -> Result<Vec<RgbaImage>> {
|
pub fn inference(&self, mes: ChatCompletionParameters) -> Result<Vec<RgbaImage>> {
|
||||||
let imgs = extract_images(&mes)?;
|
let imgs = extract_images(&mes)?;
|
||||||
let mut rmbg_png = vec![];
|
let mut rmbg_png = vec![];
|
||||||
for img in imgs {
|
for img in imgs {
|
||||||
@@ -78,3 +86,39 @@ impl RMBG2_0 {
|
|||||||
Ok(rmbg_png)
|
Ok(rmbg_png)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl GenerateModel for RMBG2_0Model {
|
||||||
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||||
|
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<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
||||||
|
+ Send
|
||||||
|
+ Unpin
|
||||||
|
+ '_,
|
||||||
|
>,
|
||||||
|
> {
|
||||||
|
let error_stream = stream::once(async {
|
||||||
|
Err(anyhow::anyhow!(format!(
|
||||||
|
"{} model not support stream",
|
||||||
|
self.model_name
|
||||||
|
))) as Result<ChatCompletionChunkResponse, anyhow::Error>
|
||||||
|
});
|
||||||
|
|
||||||
|
Ok(Box::new(Box::pin(error_stream)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,22 +1,36 @@
|
|||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
|
|
||||||
|
use aha_openai_dive::v1::resources::chat::{
|
||||||
|
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||||
|
};
|
||||||
use anyhow::{Ok, Result};
|
use anyhow::{Ok, Result};
|
||||||
|
use base64::{Engine, prelude::BASE64_STANDARD};
|
||||||
use candle_core::{DType, Device, Tensor, pickle::read_all_with_key};
|
use candle_core::{DType, Device, Tensor, pickle::read_all_with_key};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
|
use rocket::futures::{Stream, stream};
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::voxcpm::{
|
models::{
|
||||||
audio_vae::AudioVAE,
|
GenerateModel,
|
||||||
config::{AudioVaeConfig, VoxCPMConfig},
|
voxcpm::{
|
||||||
model::VoxCPMModel,
|
audio_vae::AudioVAE,
|
||||||
tokenizer::SingleChineseTokenizer,
|
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 {
|
pub struct VoxCPMGenerate {
|
||||||
voxcpm: VoxCPMModel,
|
voxcpm: VoxCPMModel,
|
||||||
prompt_cache: Option<HashMap<String, Tensor>>,
|
prompt_cache: Option<HashMap<String, Tensor>>,
|
||||||
|
sample_rate: usize,
|
||||||
|
model_name: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl VoxCPMGenerate {
|
impl VoxCPMGenerate {
|
||||||
@@ -48,6 +62,11 @@ impl VoxCPMGenerate {
|
|||||||
sample_rate: 16000,
|
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(
|
let audio_vae = AudioVAE::new(
|
||||||
vb_vae,
|
vb_vae,
|
||||||
audio_config.encoder_dim,
|
audio_config.encoder_dim,
|
||||||
@@ -85,6 +104,8 @@ impl VoxCPMGenerate {
|
|||||||
Ok(Self {
|
Ok(Self {
|
||||||
voxcpm,
|
voxcpm,
|
||||||
prompt_cache: None,
|
prompt_cache: None,
|
||||||
|
sample_rate: audio_config.sample_rate,
|
||||||
|
model_name,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -135,7 +156,7 @@ impl VoxCPMGenerate {
|
|||||||
prompt_text: Option<String>,
|
prompt_text: Option<String>,
|
||||||
prompt_wav_path: Option<String>,
|
prompt_wav_path: Option<String>,
|
||||||
) -> Result<Tensor> {
|
) -> Result<Tensor> {
|
||||||
let audio = self.generate(
|
let audio = self.inference(
|
||||||
target_text,
|
target_text,
|
||||||
prompt_text,
|
prompt_text,
|
||||||
prompt_wav_path,
|
prompt_wav_path,
|
||||||
@@ -150,10 +171,10 @@ impl VoxCPMGenerate {
|
|||||||
}
|
}
|
||||||
pub fn generate_simple(&mut self, target_text: String) -> Result<Tensor> {
|
pub fn generate_simple(&mut self, target_text: String) -> Result<Tensor> {
|
||||||
// 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, 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)
|
Ok(audio)
|
||||||
}
|
}
|
||||||
pub fn generate(
|
pub fn inference(
|
||||||
&mut self,
|
&mut self,
|
||||||
target_text: String,
|
target_text: String,
|
||||||
prompt_text: Option<String>,
|
prompt_text: Option<String>,
|
||||||
@@ -179,3 +200,59 @@ impl VoxCPMGenerate {
|
|||||||
Ok(audio)
|
Ok(audio)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl GenerateModel for VoxCPMGenerate {
|
||||||
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||||
|
let prompt_text = extract_metadata_value::<String>(&mes.metadata, "prompt_text");
|
||||||
|
let min_len = extract_metadata_value::<usize>(&mes.metadata, "min_len").unwrap_or(2);
|
||||||
|
let max_len = extract_metadata_value::<usize>(&mes.metadata, "max_len").unwrap_or(4096);
|
||||||
|
let inference_timesteps =
|
||||||
|
extract_metadata_value::<usize>(&mes.metadata, "inference_timesteps").unwrap_or(10);
|
||||||
|
let cfg_value = extract_metadata_value::<f64>(&mes.metadata, "cfg_value").unwrap_or(2.0);
|
||||||
|
let retry_badcase_ratio_threshold =
|
||||||
|
extract_metadata_value::<f64>(&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<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
||||||
|
+ Send
|
||||||
|
+ Unpin
|
||||||
|
+ '_,
|
||||||
|
>,
|
||||||
|
> {
|
||||||
|
let error_stream = stream::once(async {
|
||||||
|
Err(anyhow::anyhow!(format!(
|
||||||
|
"{} model not support stream",
|
||||||
|
self.model_name
|
||||||
|
))) as Result<ChatCompletionChunkResponse, anyhow::Error>
|
||||||
|
});
|
||||||
|
|
||||||
|
Ok(Box::new(Box::pin(error_stream)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -513,7 +513,7 @@ impl VoxCPMModel {
|
|||||||
let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
|
let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
|
||||||
let text_length = text_token.dim(0)?;
|
let text_length = text_token.dim(0)?;
|
||||||
let mut audio =
|
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;
|
let patch_len = self.patch_size * self.chunk_size;
|
||||||
if audio.dim(1)? % patch_len != 0 {
|
if audio.dim(1)? % patch_len != 0 {
|
||||||
audio = audio.pad_with_zeros(
|
audio = audio.pad_with_zeros(
|
||||||
@@ -728,8 +728,11 @@ impl VoxCPMModel {
|
|||||||
) -> Result<HashMap<String, Tensor>> {
|
) -> Result<HashMap<String, Tensor>> {
|
||||||
let text_token = self.tokenizer.encode(prompt_text)?;
|
let text_token = self.tokenizer.encode(prompt_text)?;
|
||||||
let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?;
|
let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?;
|
||||||
let mut audio =
|
let mut audio = load_audio_with_resample(
|
||||||
load_audio_with_resample(prompt_wav_path, self.device.clone(), Some(self.sample_rate))?;
|
&prompt_wav_path,
|
||||||
|
self.device.clone(),
|
||||||
|
Some(self.sample_rate),
|
||||||
|
)?;
|
||||||
let patch_len = self.patch_size * self.chunk_size;
|
let patch_len = self.patch_size * self.chunk_size;
|
||||||
if audio.dim(1)? % patch_len != 0 {
|
if audio.dim(1)? % patch_len != 0 {
|
||||||
audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?;
|
audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?;
|
||||||
|
|||||||
+180
-6
@@ -1,12 +1,22 @@
|
|||||||
use std::f64::consts::PI;
|
use std::fs::File;
|
||||||
use std::path::Path;
|
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 anyhow::{Result, anyhow};
|
||||||
|
use base64::Engine;
|
||||||
|
use base64::prelude::BASE64_STANDARD;
|
||||||
use candle_core::{D, Device, Tensor};
|
use candle_core::{D, Device, Tensor};
|
||||||
use candle_nn::{Conv1d, Conv1dConfig, Module};
|
use candle_nn::{Conv1d, Conv1dConfig, Module};
|
||||||
use hound::{SampleFormat, WavReader};
|
use hound::{SampleFormat, WavReader};
|
||||||
use num::integer::gcd;
|
use num::integer::gcd;
|
||||||
|
|
||||||
|
use crate::utils::get_default_save_dir;
|
||||||
|
|
||||||
// 重采样方法枚举
|
// 重采样方法枚举
|
||||||
#[derive(Debug, Clone, Copy)]
|
#[derive(Debug, Clone, Copy)]
|
||||||
pub enum ResamplingMethod {
|
pub enum ResamplingMethod {
|
||||||
@@ -213,9 +223,62 @@ pub fn resample_simple(waveform: &Tensor, orig_freq: i64, new_freq: i64) -> Resu
|
|||||||
None,
|
None,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
pub fn load_audio_from_url(url: &str) -> Result<PathBuf> {
|
||||||
|
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<P: AsRef<Path>>(path: P, device: Device) -> Result<(Tensor, usize)> {
|
let mut file = std::fs::File::create(&temp_path)?;
|
||||||
let mut reader = WavReader::open(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<PathBuf> {
|
||||||
|
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 spec = reader.spec();
|
||||||
let samples: Vec<f32> = match spec.sample_format {
|
let samples: Vec<f32> = match spec.sample_format {
|
||||||
SampleFormat::Int => {
|
SampleFormat::Int => {
|
||||||
@@ -264,8 +327,8 @@ pub fn load_audio<P: AsRef<Path>>(path: P, device: Device) -> Result<(Tensor, us
|
|||||||
Ok((audio_tensor, sample_rate as usize))
|
Ok((audio_tensor, sample_rate as usize))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn load_audio_with_resample<P: AsRef<Path>>(
|
pub fn load_audio_with_resample(
|
||||||
path: P,
|
path: &str,
|
||||||
device: Device,
|
device: Device,
|
||||||
target_sample_rate: Option<usize>,
|
target_sample_rate: Option<usize>,
|
||||||
) -> Result<Tensor> {
|
) -> Result<Tensor> {
|
||||||
@@ -299,3 +362,114 @@ pub fn save_wav(audio: &Tensor, save_path: &str, sample_rate: u32) -> Result<()>
|
|||||||
writer.finalize().unwrap();
|
writer.finalize().unwrap();
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn get_audio_wav_u8(audio: &Tensor, sample_rate: u32) -> Result<Vec<u8>> {
|
||||||
|
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::<f32>()?;
|
||||||
|
let ratio = if max > 1.0 { 32767.0 / max } else { 32767.0 };
|
||||||
|
let audio = audio.squeeze(0)?;
|
||||||
|
let audio_vec = audio.to_vec1::<f32>()?;
|
||||||
|
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<Vec<String>> {
|
||||||
|
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<Vec<String>> {
|
||||||
|
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<P: AsRef<Path>>(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<Vec<String>> {
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
|||||||
+136
-8
@@ -3,19 +3,21 @@ pub mod img_utils;
|
|||||||
pub mod tensor_utils;
|
pub mod tensor_utils;
|
||||||
pub mod video_utils;
|
pub mod video_utils;
|
||||||
|
|
||||||
use std::process::Command;
|
use std::{fs, process::Command};
|
||||||
|
|
||||||
use aha_openai_dive::v1::resources::{
|
use aha_openai_dive::v1::resources::{
|
||||||
chat::{
|
chat::{
|
||||||
ChatCompletionChoice, ChatCompletionChunkChoice, ChatCompletionChunkResponse,
|
AudioUrlType, ChatCompletionChoice, ChatCompletionChunkChoice, ChatCompletionChunkResponse,
|
||||||
ChatCompletionParameters, ChatCompletionResponse, ChatMessage, ChatMessageContent,
|
ChatCompletionParameters, ChatCompletionResponse, ChatMessage, ChatMessageAudioContentPart,
|
||||||
ChatMessageContentPart, DeltaChatMessage, DeltaFunction, DeltaToolCall, Function, ToolCall,
|
ChatMessageContent, ChatMessageContentPart, ChatMessageImageContentPart, DeltaChatMessage,
|
||||||
|
DeltaFunction, DeltaToolCall, Function, ImageUrlType, ToolCall,
|
||||||
},
|
},
|
||||||
shared::{FinishReason, Usage},
|
shared::{FinishReason, Usage},
|
||||||
};
|
};
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device};
|
use candle_core::{DType, Device};
|
||||||
use candle_transformers::generation::{LogitsProcessor, Sampling};
|
use candle_transformers::generation::{LogitsProcessor, Sampling};
|
||||||
|
use dirs::home_dir;
|
||||||
|
|
||||||
pub fn get_device(device: Option<&Device>) -> Device {
|
pub fn get_device(device: Option<&Device>) -> Device {
|
||||||
match device {
|
match device {
|
||||||
@@ -137,6 +139,90 @@ pub fn ceil_by_factor(num: f32, factor: u32) -> u32 {
|
|||||||
ceil * factor
|
ceil * factor
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn build_img_completion_response(
|
||||||
|
base64vec: &Vec<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 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(
|
pub fn build_completion_response(
|
||||||
res: String,
|
res: String,
|
||||||
model_name: &str,
|
model_name: &str,
|
||||||
@@ -144,10 +230,6 @@ pub fn build_completion_response(
|
|||||||
) -> ChatCompletionResponse {
|
) -> ChatCompletionResponse {
|
||||||
let id = uuid::Uuid::new_v4().to_string();
|
let id = uuid::Uuid::new_v4().to_string();
|
||||||
let usage = num_tokens.map(|num| Usage {
|
let usage = num_tokens.map(|num| Usage {
|
||||||
input_tokens: None,
|
|
||||||
input_tokens_details: None,
|
|
||||||
output_tokens: None,
|
|
||||||
output_tokens_details: None,
|
|
||||||
prompt_tokens: None,
|
prompt_tokens: None,
|
||||||
completion_tokens: None,
|
completion_tokens: None,
|
||||||
total_tokens: num,
|
total_tokens: num,
|
||||||
@@ -358,3 +440,49 @@ pub fn extract_mes(mes: &ChatCompletionParameters) -> Result<Vec<(String, String
|
|||||||
}
|
}
|
||||||
Ok(mes_vec)
|
Ok(mes_vec)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn extract_metadata_value<T>(
|
||||||
|
metadata: &Option<std::collections::HashMap<String, String>>,
|
||||||
|
key: &str,
|
||||||
|
) -> Option<T>
|
||||||
|
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::<T>()
|
||||||
|
{
|
||||||
|
return Some(value);
|
||||||
|
}
|
||||||
|
None
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn extract_user_text(mes: &ChatCompletionParameters) -> Result<String> {
|
||||||
|
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<String> {
|
||||||
|
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()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|||||||
+12
-10
@@ -1,3 +1,4 @@
|
|||||||
|
use aha::utils::get_default_save_dir;
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::Tensor;
|
use candle_core::Tensor;
|
||||||
|
|
||||||
@@ -5,18 +6,19 @@ use candle_core::Tensor;
|
|||||||
fn messy_test() -> Result<()> {
|
fn messy_test() -> Result<()> {
|
||||||
// RUST_BACKTRACE=1 cargo test -F cuda messy_test -r -- --nocapture
|
// RUST_BACKTRACE=1 cargo test -F cuda messy_test -r -- --nocapture
|
||||||
let device = &candle_core::Device::Cpu;
|
let device = &candle_core::Device::Cpu;
|
||||||
|
// let path = get_default_save_dir();
|
||||||
let x = Tensor::arange(0.0, 9.0, device)?;
|
let x = Tensor::arange(0.0, 9.0, device)?;
|
||||||
println!("x: {}", x);
|
println!("x: {}", x);
|
||||||
let x = x
|
// let x = x
|
||||||
.unsqueeze(0)?
|
// .unsqueeze(0)?
|
||||||
.unsqueeze(0)?
|
// .unsqueeze(0)?
|
||||||
.broadcast_as((5, 5, 9))?
|
// .broadcast_as((5, 5, 9))?
|
||||||
.reshape((5, 5, 3, 3))?;
|
// .reshape((5, 5, 3, 3))?;
|
||||||
println!("x: {}", x);
|
// println!("x: {}", x);
|
||||||
let x = x.permute((0, 2, 1, 3))?;
|
// let x = x.permute((0, 2, 1, 3))?;
|
||||||
println!("x: {}", x);
|
// println!("x: {}", x);
|
||||||
let x = x.reshape((15, 15))?;
|
// let x = x.reshape((15, 15))?;
|
||||||
println!("x: {}", x);
|
// println!("x: {}", x);
|
||||||
// let xs = Tensor::rand(0.0, 5.0, (1, 1, 3, 3), device)?;
|
// let xs = Tensor::rand(0.0, 5.0, (1, 1, 3, 3), device)?;
|
||||||
// println!("xs: {}", xs);
|
// println!("xs: {}", xs);
|
||||||
// let xs = xs.pad_with_zeros(3, 2, 2)?
|
// let xs = xs.pad_with_zeros(3, 2, 2)?
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
use std::time::Instant;
|
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 aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
|
|
||||||
@@ -31,12 +31,12 @@ fn rmbg2_0_generate() -> Result<()> {
|
|||||||
"#;
|
"#;
|
||||||
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
||||||
let i_start = Instant::now();
|
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();
|
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 result = model.generate(mes)?;
|
let result = model.inference(mes)?;
|
||||||
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);
|
||||||
for (i, img) in result.iter().enumerate() {
|
for (i, img) in result.iter().enumerate() {
|
||||||
|
|||||||
+57
-7
@@ -1,14 +1,64 @@
|
|||||||
use std::time::Instant;
|
use std::time::Instant;
|
||||||
|
|
||||||
use aha::{
|
use aha::{
|
||||||
models::voxcpm::{generate::VoxCPMGenerate, tokenizer::SingleChineseTokenizer},
|
models::{
|
||||||
utils::audio_utils::save_wav,
|
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};
|
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]
|
#[test]
|
||||||
fn voxcpm_generate() -> Result<()> {
|
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 model_path = "/home/jhq/huggingface_model/openbmb/VoxCPM-0.5B/";
|
||||||
|
|
||||||
let i_start = Instant::now();
|
let i_start = Instant::now();
|
||||||
@@ -18,12 +68,12 @@ fn voxcpm_generate() -> Result<()> {
|
|||||||
|
|
||||||
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.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(),
|
"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("./assets/audio/voice_01.wav".to_string()),
|
Some("file://./assets/audio/voice_01.wav".to_string()),
|
||||||
// Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
|
// Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
|
||||||
// Some("./assets/audio/voice_05.wav".to_string()),
|
// Some("file://./assets/audio/voice_05.wav".to_string()),
|
||||||
2,
|
2,
|
||||||
100,
|
100,
|
||||||
10,
|
10,
|
||||||
@@ -35,7 +85,7 @@ fn voxcpm_generate() -> Result<()> {
|
|||||||
// 创建prompt_cache
|
// 创建prompt_cache
|
||||||
// let _ = voxcpm_generate.build_prompt_cache(
|
// let _ = voxcpm_generate.build_prompt_cache(
|
||||||
// "啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
|
// "啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
|
||||||
// "./assets/audio/voice_01.wav".to_string(),
|
// "file://./assets/audio/voice_01.wav".to_string(),
|
||||||
// )?;
|
// )?;
|
||||||
// // 使用prompt_cache生成语音
|
// // 使用prompt_cache生成语音
|
||||||
// let generate = voxcpm_generate.generate_use_prompt_cache(
|
// let generate = voxcpm_generate.generate_use_prompt_cache(
|
||||||
|
|||||||
+55
-6
@@ -1,11 +1,60 @@
|
|||||||
use std::time::Instant;
|
use std::time::Instant;
|
||||||
|
|
||||||
use aha::{
|
use aha::{
|
||||||
models::voxcpm::{generate::VoxCPMGenerate, tokenizer::SingleChineseTokenizer},
|
models::{
|
||||||
utils::audio_utils::save_wav,
|
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};
|
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]
|
#[test]
|
||||||
fn voxcpm1_5_generate() -> Result<()> {
|
fn voxcpm1_5_generate() -> Result<()> {
|
||||||
// RUST_BACKTRACE=1 cargo test -F cuda voxcpm1_5_generate -r -- --nocapture
|
// 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 i_start = Instant::now();
|
||||||
// let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?;
|
// 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(),
|
"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("./assets/audio/voice_01.wav".to_string()),
|
Some("file://./assets/audio/voice_01.wav".to_string()),
|
||||||
// Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
|
// Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
|
||||||
// Some("./assets/audio/voice_05.wav".to_string()),
|
// Some("file://./assets/audio/voice_05.wav".to_string()),
|
||||||
2,
|
2,
|
||||||
4096,
|
4096,
|
||||||
10,
|
10,
|
||||||
@@ -35,7 +84,7 @@ fn voxcpm1_5_generate() -> Result<()> {
|
|||||||
// 创建prompt_cache
|
// 创建prompt_cache
|
||||||
// let _ = voxcpm_generate.build_prompt_cache(
|
// let _ = voxcpm_generate.build_prompt_cache(
|
||||||
// "啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
|
// "啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
|
||||||
// "./assets/audio/voice_01.wav".to_string(),
|
// "file://./assets/audio/voice_01.wav".to_string(),
|
||||||
// )?;
|
// )?;
|
||||||
// // 使用prompt_cache生成语音
|
// // 使用prompt_cache生成语音
|
||||||
// let generate = voxcpm_generate.generate_use_prompt_cache(
|
// let generate = voxcpm_generate.generate_use_prompt_cache(
|
||||||
|
|||||||
Reference in New Issue
Block a user