add image api and audio api
This commit is contained in:
+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 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<String> {
|
||||
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(())
|
||||
|
||||
+26
-1
@@ -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<PaddleOCRVLGenerateModel<'a>>),
|
||||
RMBG2_0(Box<RMBG2_0Model>),
|
||||
VoxCPM(Box<VoxCPMGenerate>),
|
||||
}
|
||||
|
||||
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<ModelInstance<'_
|
||||
let model = PaddleOCRVLGenerateModel::init(path, None, None)?;
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -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<DType>) -> Result<Self> {
|
||||
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<Vec<RgbaImage>> {
|
||||
pub fn inference(&self, mes: ChatCompletionParameters) -> Result<Vec<RgbaImage>> {
|
||||
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<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 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<HashMap<String, Tensor>>,
|
||||
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<String>,
|
||||
prompt_wav_path: Option<String>,
|
||||
) -> Result<Tensor> {
|
||||
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<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, 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<String>,
|
||||
@@ -179,3 +200,59 @@ impl VoxCPMGenerate {
|
||||
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_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<HashMap<String, Tensor>> {
|
||||
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)?;
|
||||
|
||||
+180
-6
@@ -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<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 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<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 samples: Vec<f32> = match spec.sample_format {
|
||||
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))
|
||||
}
|
||||
|
||||
pub fn load_audio_with_resample<P: AsRef<Path>>(
|
||||
path: P,
|
||||
pub fn load_audio_with_resample(
|
||||
path: &str,
|
||||
device: Device,
|
||||
target_sample_rate: Option<usize>,
|
||||
) -> Result<Tensor> {
|
||||
@@ -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<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 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<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(
|
||||
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<Vec<(String, String
|
||||
}
|
||||
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()
|
||||
})
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user