add image api and audio api

This commit is contained in:
jhqxxx
2025-12-25 20:25:52 +08:00
parent bf53d3d869
commit 9a4fa636be
16 changed files with 694 additions and 72 deletions
+26 -1
View File
@@ -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)
}
+50 -6
View File
@@ -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)))
}
}
+86 -9
View File
@@ -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)))
}
}
+6 -3
View File
@@ -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)?;