diff --git a/Cargo.lock b/Cargo.lock index 70850f9..5fe9161 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -21,7 +21,7 @@ dependencies = [ [[package]] name = "aha" -version = "0.2.5" +version = "0.2.6" dependencies = [ "ahash", "anyhow", diff --git a/Cargo.toml b/Cargo.toml index 2033136..5f17556 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,10 +1,10 @@ [package] name = "aha" -version = "0.2.5" +version = "0.2.6" edition = "2024" repository = "https://github.com/jhqxxx/aha" license = "Apache-2.0" -description = "aha model inference library, now supports Qwen(2.5VL/3/3VL/3.5/ASR/3Embedding/3Reranker), MiniCPM4, VoxCPM(0.5B/1.5/2), DeepSeek-OCR/2, Hunyuan-OCR, PaddleOCR-VL/1.5, RMBG2.0, GLM(ASR-Nano-2512/OCR), Fun-ASR-Nano-2512, LFM(2/2.5/2VL/2.5VL)" +description = "aha model inference library, now supports Qwen(2.5VL/3/3VL/3.5/ASR/3Embedding/3Reranker), MiniCPM(4/5), VoxCPM(0.5B/1.5/2), DeepSeek-OCR/2, Hunyuan-OCR, PaddleOCR-VL/1.5, RMBG2.0, GLM(ASR-Nano-2512/OCR), Fun-ASR-Nano-2512, LFM(2/2.5/2VL/2.5VL)" [dependencies] candle-core = { version = "0.9.2" } @@ -21,9 +21,7 @@ base64 = "0.22.1" num = "0.4.3" minijinja = "2.12.0" tokenizers = "0.22.1" -# aha_openai_dive = { version = "1.4", features = ["stream"] } uuid = { version = "1.18.1", features = ["v4"] } -# chrono = "0.4" rocket = { version = "0.5.1", features = ["serde_json", "json"] } tokio = "1.47.1" hound = "3.5.1" diff --git a/README.md b/README.md index 9893ee3..b0401c8 100644 --- a/README.md +++ b/README.md @@ -39,6 +39,9 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an | **Reranker** | Qwen3-Reranker | ## Changelog +### 2026-05-28 +- generate code refactoring progress 1/3 + ### 2026-05-27 - add MiniCPM5 @@ -54,11 +57,6 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an ### 2026-04-25 - VoxCPM update stream -### 2026-04-17 -- Qwen3ASR add vad data recognition - -### 2026-04-16 -- fix FireRedVAD fsmn cache bug **[View full changelog](docs/changelog.md)** → diff --git a/README.zh-CN.md b/README.zh-CN.md index 380b5bd..7dd64ee 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -38,6 +38,9 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理 | **重排序** | Qwen3-Reranker | ## 更新日志 +### 2026-05-28 +- generate代码重构进度 1/3 + ### 2026-05-27 - 新增 MiniCPM5 @@ -53,11 +56,6 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理 ### 2026-04-25 - VoxCPM 更新流式生成 -### 2026-04-17 -- Qwen3ASR 增加 vad 数据识别 - -### 2026-04-16 -- 修复 FireRedVAD fsmn 缓存问题 **[查看完整更新日志](docs/changelog.zh-CN.md)** → diff --git a/docs/changelog.md b/docs/changelog.md index 9721da2..38b07b0 100644 --- a/docs/changelog.md +++ b/docs/changelog.md @@ -5,6 +5,9 @@ All notable changes to aha will be documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). +### 2026-05-28 +- generate code refactoring progress 1/3 + ### 2026-05-27 - add MiniCPM5 diff --git a/docs/changelog.zh-CN.md b/docs/changelog.zh-CN.md index 51f3088..a4b0279 100644 --- a/docs/changelog.zh-CN.md +++ b/docs/changelog.zh-CN.md @@ -5,6 +5,9 @@ 格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/), 本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。 +### 2026-05-28 +- generate代码重构进度 1/3 + ### 2026-05-27 - 新增 MiniCPM5 diff --git a/src/models/common/generate.rs b/src/models/common/generate.rs index 7e1c1ed..5201bd0 100644 --- a/src/models/common/generate.rs +++ b/src/models/common/generate.rs @@ -10,7 +10,7 @@ use crate::{ InferenceModel, MultiModalData, sample::{get_logit_processor, use_repeat_penalty}, }, - params::chat::{ChatCompletionChunkResponse, ChatCompletionResponse}, + params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, tokenizer::TokenizerModel, utils::response_utils::{ build_chunk_response_with_reasoning, build_chunk_response_with_usage, @@ -366,3 +366,116 @@ pub fn generate_stream_generic( }; Ok(stream) } + +pub struct PrepareData { + pub in_reasoning: bool, + pub input_ids: Tensor, + pub multi_model_data: MultiModalData, +} + +pub trait GenerationDataProvider { + fn get_temperature(&self, req_temp: Option) -> Option { + req_temp + } + + fn get_top_p(&self, req_top_p: Option) -> Option { + req_top_p + } + + fn get_top_k(&self, top_k: Option) -> Option { + top_k + } + + fn is_in_reasoning(&self, text: &str) -> bool { + text.ends_with("\n") + } + + fn get_multi_model_data(&self) -> MultiModalData { + MultiModalData::new(vec![]) + } + + fn get_data(&self, mes: &ChatCompletionParameters) -> Result; +} + +#[macro_export] +macro_rules! impl_generate_model { + ($struct_name: ty) => { + impl<'a> $crate::models::GenerateModel for $struct_name { + fn generate( + &mut self, + mes: $crate::params::chat::ChatCompletionParameters, + ) -> anyhow::Result<$crate::params::chat::ChatCompletionResponse> { + let seed = mes.seed.unwrap_or(299792458) as u64; + let sample_len = mes.max_tokens.unwrap_or(1024); + let temperature = self.get_temperature(mes.temperature); + let top_p = self.get_top_p(mes.top_p); + let top_k = self.get_top_k(mes.top_k); + let prepare_data = self.get_data(&mes)?; + let input_ids = prepare_data.input_ids; + let data = prepare_data.multi_model_data; + let mut ctx = $crate::models::common::generate::GenerationContext::new( + temperature, + top_p, + top_k, + mes.repeat_penalty, + mes.repeat_last_n, + seed, + input_ids.dim(1)?, + sample_len, + self.device.clone(), + ); + + $crate::models::common::generate::generate_generic( + &mut self.model, + &self.tokenizer, + input_ids, + data, + &mut ctx, + &self.model_name, + ) + } + + fn generate_stream( + &mut self, + mes: $crate::params::chat::ChatCompletionParameters, + ) -> anyhow::Result< + Box< + dyn rocket::futures::Stream< + Item = anyhow::Result< + $crate::params::chat::ChatCompletionChunkResponse, + >, + > + Send + + Unpin + + '_, + >, + > { + let seed = mes.seed.unwrap_or(299792458) as u64; + let prepare_data = self.get_data(&mes)?; + let input_ids = prepare_data.input_ids; + let data = prepare_data.multi_model_data; + let in_reasoning = prepare_data.in_reasoning; + let sample_len = mes.max_tokens.unwrap_or(1024); + let temperature = self.get_temperature(mes.temperature); + let top_p = self.get_top_p(mes.top_p); + let top_k = self.get_top_k(mes.top_k); + let stream = $crate::models::common::generate::generate_stream_generic( + &mut self.model, + &self.tokenizer, + input_ids, + data, + temperature, + top_p, + top_k, + mes.repeat_penalty, + mes.repeat_last_n, + seed, + sample_len, + in_reasoning, + &self.device, + &self.model_name, + )?; + Ok(Box::new(Box::pin(stream))) + } + } + }; +} diff --git a/src/models/deepseek_ocr/generate.rs b/src/models/deepseek_ocr/generate.rs index 9427684..a2883b7 100644 --- a/src/models/deepseek_ocr/generate.rs +++ b/src/models/deepseek_ocr/generate.rs @@ -1,21 +1,14 @@ -use crate::{ - models::common::{ - MultiModalData, - generate::{GenerationContext, generate_generic, generate_stream_generic}, - }, - params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, +use crate::models::common::{ + MultiModalData, + generate::{GenerationDataProvider, PrepareData}, }; use anyhow::Result; use candle_core::{DType, Device}; use candle_nn::VarBuilder; -use rocket::futures::Stream; use crate::{ - models::{ - GenerateModel, - deepseek_ocr::{ - config::DeepseekOCRConfig, model::DeepseekOCRModel, processor::DeepseekOCRProcessor, - }, + models::deepseek_ocr::{ + config::DeepseekOCRConfig, model::DeepseekOCRModel, processor::DeepseekOCRProcessor, }, tokenizer::TokenizerModel, utils::{extract_metadata_value, find_type_files, get_device, get_dtype}, @@ -24,9 +17,7 @@ use crate::{ pub struct DeepseekOCRGenerateModel { tokenizer: TokenizerModel, processor: DeepseekOCRProcessor, - deepseekocr_model: DeepseekOCRModel, - // bos_token_id: u32, - // eos_token_id: u32, + model: DeepseekOCRModel, device: Device, size: Vec, model_name: String, @@ -52,19 +43,15 @@ impl DeepseekOCRGenerateModel { 1usize }; let processor = DeepseekOCRProcessor::new(device, dtype, version)?; - // let eos_token_id = cfg.eos_token_id; - // let bos_token_id = cfg.bos_token_id; let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; - let deepseekocr_model = DeepseekOCRModel::new(vb, cfg, version)?; + let model = DeepseekOCRModel::new(vb, cfg, version)?; let size = vec![512u32, 640, 1024, 1280]; Ok(Self { tokenizer, processor, - deepseekocr_model, - // bos_token_id, - // eos_token_id, + model, device: device.clone(), size, model_name: model_name.to_string(), @@ -73,8 +60,8 @@ impl DeepseekOCRGenerateModel { } } -impl GenerateModel for DeepseekOCRGenerateModel { - fn generate(&mut self, mes: ChatCompletionParameters) -> Result { +impl GenerationDataProvider for DeepseekOCRGenerateModel { + fn get_data(&self, mes: &crate::params::chat::ChatCompletionParameters) -> Result { let base_size = extract_metadata_value::(&mes.metadata, "base_size").unwrap_or(640); let base_size = if self.size.contains(&base_size) { base_size @@ -92,93 +79,20 @@ impl GenerateModel for DeepseekOCRGenerateModel { let crop_mode = extract_metadata_value::(&mes.metadata, "crop_mode").unwrap_or(false); let (input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self .processor - .process_info(&mes, &self.tokenizer, base_size, image_size, crop_mode)?; - let max_tokens = mes.max_tokens.unwrap_or(1024); - let mut ctx = GenerationContext::new( - mes.temperature, - mes.top_p, - None, - mes.repeat_penalty, - mes.repeat_last_n, - mes.seed.unwrap_or(34562) as u64, - input_ids.dim(1)?, - max_tokens, - self.device.clone(), - ); + .process_info(mes, &self.tokenizer, base_size, image_size, crop_mode)?; let data_vec = vec![ Some(images_ori), Some(image_crop), Some(images_seq_mask), Some(images_spatial_crop_t), ]; - let data = MultiModalData::new(data_vec); - generate_generic( - &mut self.deepseekocr_model, - &self.tokenizer, + let multi_model_data = MultiModalData::new(data_vec); + Ok(PrepareData { + in_reasoning: false, input_ids, - data, - &mut ctx, - &self.model_name, - ) - } - - fn generate_stream( - &mut self, - mes: ChatCompletionParameters, - ) -> Result< - Box< - dyn Stream> - + Send - + Unpin - + '_, - >, - > { - let base_size = extract_metadata_value::(&mes.metadata, "base_size").unwrap_or(640); - let base_size = if self.size.contains(&base_size) { - base_size - } else { - 640 - }; - let image_size = extract_metadata_value::(&mes.metadata, "image_size").unwrap_or(640); - let image_size = if self.size.contains(&image_size) { - image_size - } else { - 640 - }; - let base_size = if self.version == 2 { 1024 } else { base_size }; - let image_size = if self.version == 2 { 768 } else { image_size }; - let crop_mode = extract_metadata_value::(&mes.metadata, "crop_mode").unwrap_or(false); - let (input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self - .processor - .process_info(&mes, &self.tokenizer, base_size, image_size, crop_mode)?; - let data_vec = vec![ - images_ori.into(), - image_crop.into(), - images_seq_mask.into(), - images_spatial_crop_t.into(), - ]; - let data = MultiModalData::new(data_vec); - - let temperature = mes.temperature; - let top_p = mes.top_p; - let seed = mes.seed.unwrap_or(34562) as u64; - let max_tokens = mes.max_tokens.unwrap_or(1024); - let stream = generate_stream_generic( - &mut self.deepseekocr_model, - &self.tokenizer, - input_ids, - data, - temperature, - top_p, - None, - mes.repeat_penalty, - mes.repeat_last_n, - seed, - max_tokens, - false, - &self.device, - &self.model_name, - )?; - Ok(Box::new(Box::pin(stream))) + multi_model_data, + }) } } + +crate::impl_generate_model!(DeepseekOCRGenerateModel); diff --git a/src/models/fun_asr_nano/generate.rs b/src/models/fun_asr_nano/generate.rs index e7ef42c..75a8cc2 100644 --- a/src/models/fun_asr_nano/generate.rs +++ b/src/models/fun_asr_nano/generate.rs @@ -27,7 +27,7 @@ use crate::{ pub struct FunAsrNanoGenerateModel { tokenizer: TokenizerModel, processor: FunAsrNanoProcessor, - fun_asr_nano: FunAsrNanoModel, + model: FunAsrNanoModel, device: Device, dtype: DType, generation_config: Qwen3GenerationConfig, @@ -75,7 +75,7 @@ impl FunAsrNanoGenerateModel { } } let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &device); - let fun_asr_nano = + let model = FunAsrNanoModel::new(vb, &cfg, &llm_cfg, generation_config.eos_token_id.clone())?; let model_name = std::path::Path::new(path) .file_name() @@ -85,7 +85,7 @@ impl FunAsrNanoGenerateModel { Ok(Self { tokenizer, processor, - fun_asr_nano, + model, device, dtype, generation_config, @@ -120,7 +120,7 @@ impl GenerateModel for FunAsrNanoGenerateModel { let data_vec = vec![speech.into(), fbank_mask.into()]; let data = MultiModalData::new(data_vec); generate_generic( - &mut self.fun_asr_nano, + &mut self.model, &self.tokenizer, input_ids, data, @@ -152,7 +152,7 @@ impl GenerateModel for FunAsrNanoGenerateModel { let data_vec = vec![speech.into(), fbank_mask.into()]; let data = MultiModalData::new(data_vec); let stream = generate_stream_generic( - &mut self.fun_asr_nano, + &mut self.model, &self.tokenizer, input_ids, data, diff --git a/src/models/glm_asr_nano/generate.rs b/src/models/glm_asr_nano/generate.rs index 117239d..81d3a17 100644 --- a/src/models/glm_asr_nano/generate.rs +++ b/src/models/glm_asr_nano/generate.rs @@ -26,7 +26,7 @@ pub struct GlmAsrNanoGenerateModel<'a> { chat_template: ChatTemplate<'a>, tokenizer: TokenizerModel, processor: GlmAsrNanoProcessor, - glm_asr_nano: GlmAsrNanoModel, + model: GlmAsrNanoModel, device: Device, dtype: DType, model_name: String, @@ -45,7 +45,7 @@ impl<'a> GlmAsrNanoGenerateModel<'a> { let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; let eos_ids = vec![59246u32, 59253, 59255]; - let glm_asr_nano = GlmAsrNanoModel::new(vb, cfg, eos_ids)?; + let model = GlmAsrNanoModel::new(vb, cfg, eos_ids)?; let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -55,7 +55,7 @@ impl<'a> GlmAsrNanoGenerateModel<'a> { chat_template, tokenizer, processor, - glm_asr_nano, + model, device, dtype, model_name, @@ -88,7 +88,7 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> { let data_vec = vec![input_features.into(), audio_token_lengths.into()]; let data = MultiModalData::new(data_vec); generate_generic( - &mut self.glm_asr_nano, + &mut self.model, &self.tokenizer, input_ids, data, @@ -119,7 +119,7 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> { let data_vec = vec![input_features.into(), audio_token_lengths.into()]; let data = MultiModalData::new(data_vec); let stream = generate_stream_generic( - &mut self.glm_asr_nano, + &mut self.model, &self.tokenizer, input_ids, data, diff --git a/src/models/hunyuan_ocr/generate.rs b/src/models/hunyuan_ocr/generate.rs index df5d7a4..7e49f84 100644 --- a/src/models/hunyuan_ocr/generate.rs +++ b/src/models/hunyuan_ocr/generate.rs @@ -28,7 +28,7 @@ pub struct HunyuanOCRGenerateModel<'a> { chat_template: ChatTemplate<'a>, tokenizer: TokenizerModel, pre_processor: HunyuanVLProcessor, - hunyuan_vl: HunyuanVLModel, + model: HunyuanVLModel, device: Device, generation_config: HunyuanOCRGenerationConfig, model_name: String, @@ -49,8 +49,7 @@ impl<'a> HunyuanOCRGenerateModel<'a> { let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config: HunyuanOCRGenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; - let hunyuan_vl = - HunyuanVLModel::new(vb, cfg.clone(), generation_config.eos_token_id.clone())?; + let model = HunyuanVLModel::new(vb, cfg.clone(), generation_config.eos_token_id.clone())?; let model_name = std::path::Path::new(path) .file_name() @@ -61,7 +60,7 @@ impl<'a> HunyuanOCRGenerateModel<'a> { chat_template, tokenizer, pre_processor, - hunyuan_vl, + model, device, generation_config, model_name, @@ -103,7 +102,7 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> { ]; let data = MultiModalData::new(data_vec); generate_generic( - &mut self.hunyuan_vl, + &mut self.model, &self.tokenizer, input_ids, data, @@ -144,7 +143,7 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> { ]; let data = MultiModalData::new(data_vec); let stream = generate_stream_generic( - &mut self.hunyuan_vl, + &mut self.model, &self.tokenizer, input_ids, data, diff --git a/src/models/lfm2/generate.rs b/src/models/lfm2/generate.rs index 3020a3a..a61b7c6 100644 --- a/src/models/lfm2/generate.rs +++ b/src/models/lfm2/generate.rs @@ -1,16 +1,9 @@ -use crate::models::common::MultiModalData; -use crate::models::common::generate::{ - GenerationContext, generate_generic, generate_stream_generic, -}; -use crate::params::chat::{ChatCompletionParameters, ChatCompletionResponse}; +use crate::models::common::generate::{GenerationDataProvider, PrepareData}; use crate::{ chat_template::ChatTemplate, - models::{ - GenerateModel, - lfm2::{ - config::{Lfm2Config, Lfm2GenerateConfig}, - model::Lfm2Model, - }, + models::lfm2::{ + config::{Lfm2Config, Lfm2GenerateConfig}, + model::Lfm2Model, }, tokenizer::TokenizerModel, utils::{find_type_files, get_device, get_dtype}, @@ -62,68 +55,18 @@ impl<'a> Lfm2GenerateModel<'a> { } } -impl<'a> GenerateModel for Lfm2GenerateModel<'a> { - fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let mes_render = self.chat_template.apply_chat_template(&mes)?; +impl<'a> GenerationDataProvider for Lfm2GenerateModel<'a> { + fn get_data(&self, mes: &crate::params::chat::ChatCompletionParameters) -> Result { + let mes_render = self.chat_template.apply_chat_template(mes)?; + let in_reasoning = self.is_in_reasoning(&mes_render); let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let sample_len = mes.max_tokens.unwrap_or(1024); - let seed = mes.seed.unwrap_or(34562) as u64; - let mut ctx = GenerationContext::new( - mes.temperature, - mes.top_p, - None, - mes.repeat_penalty, - mes.repeat_last_n, - seed, - input_ids.dim(1)?, - sample_len, - self.device.clone(), - ); - - let data = MultiModalData::new(vec![]); - generate_generic( - &mut self.model, - &self.tokenizer, + let multi_model_data = self.get_multi_model_data(); + Ok(PrepareData { + in_reasoning, input_ids, - data, - &mut ctx, - &self.model_name, - ) - } - - fn generate_stream( - &mut self, - mes: ChatCompletionParameters, - ) -> Result< - Box< - dyn rocket::futures::Stream< - Item = Result, - > + Send - + Unpin - + '_, - >, - > { - let mes_render = self.chat_template.apply_chat_template(&mes)?; - let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let sample_len = mes.max_tokens.unwrap_or(1024); - let data = MultiModalData::new(vec![]); - let seed = mes.seed.unwrap_or(34562) as u64; - let stream = generate_stream_generic( - &mut self.model, - &self.tokenizer, - input_ids, - data, - mes.temperature, - mes.top_p, - None, - mes.repeat_penalty, - mes.repeat_last_n, - seed, - sample_len, - false, - &self.device, - &self.model_name, - )?; - Ok(Box::new(Box::pin(stream))) + multi_model_data, + }) } } + +crate::impl_generate_model!(Lfm2GenerateModel<'a>); diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs index de3a0ce..f99b0e7 100644 --- a/src/models/minicpm4/generate.rs +++ b/src/models/minicpm4/generate.rs @@ -1,25 +1,18 @@ -use crate::models::common::MultiModalData; -use crate::models::common::generate::{ - GenerationContext, generate_generic, generate_stream_generic, -}; -use crate::params::chat::{ - ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, -}; +use crate::models::common::generate::{GenerationDataProvider, PrepareData}; + use anyhow::Result; use candle_core::{DType, Device}; use candle_nn::VarBuilder; -use rocket::futures::Stream; use crate::models::minicpm4::config::MiniCPM4Config; use crate::models::minicpm4::model::MiniCPMModel; -// use crate::models::GenerateStream; use crate::utils::{find_type_files, get_device, get_dtype}; -use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel}; +use crate::{chat_template::ChatTemplate, tokenizer::TokenizerModel}; pub struct MiniCPM4GenerateModel<'a> { chat_template: ChatTemplate<'a>, tokenizer: TokenizerModel, - minicpm: MiniCPMModel, + model: MiniCPMModel, device: Device, model_name: String, } @@ -35,7 +28,7 @@ impl<'a> MiniCPM4GenerateModel<'a> { let dtype = get_dtype(dtype, cfg_dtype); let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; - let minicpm = MiniCPMModel::new(vb, cfg)?; + let model = MiniCPMModel::new(vb, cfg)?; let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -44,73 +37,25 @@ impl<'a> MiniCPM4GenerateModel<'a> { Ok(MiniCPM4GenerateModel { chat_template, tokenizer, - minicpm, + model, device: device.clone(), model_name, }) } } -impl<'a> GenerateModel for MiniCPM4GenerateModel<'a> { - fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let mes_render = self.chat_template.apply_chat_template(&mes)?; +impl<'a> GenerationDataProvider for MiniCPM4GenerateModel<'a> { + fn get_data(&self, mes: &crate::params::chat::ChatCompletionParameters) -> Result { + let mes_render = self.chat_template.apply_chat_template(mes)?; + let in_reasoning = self.is_in_reasoning(&mes_render); let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let seed = mes.seed.unwrap_or(34562) as u64; - let sample_len = mes.max_tokens.unwrap_or(2048); - let mut ctx = GenerationContext::new( - mes.temperature, - mes.top_p, - None, - mes.repeat_penalty, - mes.repeat_last_n, - seed, - input_ids.dim(1)?, - sample_len, - self.device.clone(), - ); - - let data = MultiModalData::new(vec![]); - generate_generic( - &mut self.minicpm, - &self.tokenizer, + let multi_model_data = self.get_multi_model_data(); + Ok(PrepareData { + in_reasoning, input_ids, - data, - &mut ctx, - &self.model_name, - ) - } - fn generate_stream( - &mut self, - mes: ChatCompletionParameters, - ) -> Result< - Box< - dyn Stream> - + Send - + Unpin - + '_, - >, - > { - let seed = mes.seed.unwrap_or(34562) as u64; - let mes_render = self.chat_template.apply_chat_template(&mes)?; - let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let data = MultiModalData::new(vec![]); - let sample_len = mes.max_tokens.unwrap_or(512); - let stream = generate_stream_generic( - &mut self.minicpm, - &self.tokenizer, - input_ids, - data, - mes.temperature, - mes.top_p, - None, - mes.repeat_penalty, - mes.repeat_last_n, - seed, - sample_len, - false, - &self.device, - &self.model_name, - )?; - Ok(Box::new(Box::pin(stream))) + multi_model_data, + }) } } + +crate::impl_generate_model!(MiniCPM4GenerateModel<'a>); diff --git a/src/models/minicpm5/generate.rs b/src/models/minicpm5/generate.rs index eaf14e8..edc595a 100644 --- a/src/models/minicpm5/generate.rs +++ b/src/models/minicpm5/generate.rs @@ -1,19 +1,13 @@ -use crate::models::common::MultiModalData; -use crate::models::common::generate::{ - GenerationContext, generate_generic, generate_stream_generic, -}; +use crate::models::common::generate::{GenerationDataProvider, PrepareData}; use crate::models::llama::LlamaForCausalLM; use crate::models::minicpm5::config::MiniCPM5Config; -use crate::params::chat::{ - ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, -}; + use anyhow::Result; use candle_core::{DType, Device}; use candle_nn::VarBuilder; -use rocket::futures::Stream; use crate::utils::{find_type_files, get_device, get_dtype}; -use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel}; +use crate::{chat_template::ChatTemplate, tokenizer::TokenizerModel}; pub struct MiniCPM5GenerateModel<'a> { chat_template: ChatTemplate<'a>, @@ -70,66 +64,18 @@ impl<'a> MiniCPM5GenerateModel<'a> { } } -impl<'a> GenerateModel for MiniCPM5GenerateModel<'a> { - fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let mes_render = self.chat_template.apply_chat_template(&mes)?; +impl<'a> GenerationDataProvider for MiniCPM5GenerateModel<'a> { + fn get_data(&self, mes: &crate::params::chat::ChatCompletionParameters) -> Result { + let mes_render = self.chat_template.apply_chat_template(mes)?; + let in_reasoning = self.is_in_reasoning(&mes_render); let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let seed = mes.seed.unwrap_or(34562) as u64; - let sample_len = mes.max_tokens.unwrap_or(2048); - let mut ctx = GenerationContext::new( - mes.temperature, - mes.top_p, - None, - mes.repeat_penalty, - mes.repeat_last_n, - seed, - input_ids.dim(1)?, - sample_len, - self.device.clone(), - ); - - let data = MultiModalData::new(vec![]); - generate_generic( - &mut self.model, - &self.tokenizer, + let multi_model_data = self.get_multi_model_data(); + Ok(PrepareData { + in_reasoning, input_ids, - data, - &mut ctx, - &self.model_name, - ) - } - fn generate_stream( - &mut self, - mes: ChatCompletionParameters, - ) -> Result< - Box< - dyn Stream> - + Send - + Unpin - + '_, - >, - > { - let seed = mes.seed.unwrap_or(34562) as u64; - let mes_render = self.chat_template.apply_chat_template(&mes)?; - let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let data = MultiModalData::new(vec![]); - let sample_len = mes.max_tokens.unwrap_or(512); - let stream = generate_stream_generic( - &mut self.model, - &self.tokenizer, - input_ids, - data, - mes.temperature, - mes.top_p, - None, - mes.repeat_penalty, - mes.repeat_last_n, - seed, - sample_len, - false, - &self.device, - &self.model_name, - )?; - Ok(Box::new(Box::pin(stream))) + multi_model_data, + }) } } + +crate::impl_generate_model!(MiniCPM5GenerateModel<'a>); diff --git a/src/models/paddleocr_vl/generate.rs b/src/models/paddleocr_vl/generate.rs index 0c8cd87..a2a6a29 100644 --- a/src/models/paddleocr_vl/generate.rs +++ b/src/models/paddleocr_vl/generate.rs @@ -21,7 +21,7 @@ pub struct PaddleOCRVLGenerateModel<'a> { chat_template: ChatTemplate<'a>, tokenizer: TokenizerModel, pre_processor: PaddleOCRVLProcessor, - paddleocr_vl: PaddleOCRVLModel, + model: PaddleOCRVLModel, cfg: PaddleOCRVLConfig, device: Device, model_name: String, @@ -42,7 +42,7 @@ impl<'a> PaddleOCRVLGenerateModel<'a> { let pre_processor = PaddleOCRVLProcessor::new(processor_cfg, device, dtype)?; let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; - let paddleocr_vl = PaddleOCRVLModel::new(cfg.clone(), vb, vec![2])?; + let model = PaddleOCRVLModel::new(cfg.clone(), vb, vec![2])?; let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -52,7 +52,7 @@ impl<'a> PaddleOCRVLGenerateModel<'a> { chat_template, tokenizer, pre_processor, - paddleocr_vl, + model, cfg, device: device.clone(), model_name, @@ -94,7 +94,7 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> { ]; let data = MultiModalData::new(data_vec); generate_generic( - &mut self.paddleocr_vl, + &mut self.model, &self.tokenizer, input_ids, data, @@ -136,7 +136,7 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> { let data = MultiModalData::new(data_vec); let seed = mes.seed.unwrap_or(34562) as u64; let stream = generate_stream_generic( - &mut self.paddleocr_vl, + &mut self.model, &self.tokenizer, input_ids, data, diff --git a/src/models/qwen2_5vl/generate.rs b/src/models/qwen2_5vl/generate.rs index 6ddaf1e..2e0c6bc 100644 --- a/src/models/qwen2_5vl/generate.rs +++ b/src/models/qwen2_5vl/generate.rs @@ -30,7 +30,7 @@ pub struct Qwen2_5VLGenerateModel<'a> { chat_template: ChatTemplate<'a>, tokenizer: TokenizerModel, pre_processor: Qwen2_5VLProcessor, - qwen2_5_vl: Qwen2_5VLModel, + model: Qwen2_5VLModel, device: Device, endoftext_id: u32, im_end_id: u32, @@ -52,7 +52,7 @@ impl<'a> Qwen2_5VLGenerateModel<'a> { // let model_list = find_safetensors_files(&path)?; let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; - let qwen2_5_vl = Qwen2_5VLModel::new(cfg, vb)?; + let model = Qwen2_5VLModel::new(cfg, vb)?; let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -62,7 +62,7 @@ impl<'a> Qwen2_5VLGenerateModel<'a> { chat_template, tokenizer, pre_processor, - qwen2_5_vl, + model, device: device.clone(), endoftext_id, im_end_id, @@ -102,7 +102,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { let mut completion_secs = 0.0f64; for _ in 0..sample_len { let i_start = Instant::now(); - let logits = self.qwen2_5_vl.forward( + let logits = self.model.forward( &input_ids, pixel_values, image_grid_thw, @@ -136,7 +136,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { } let num_token = generate.len() as u32; let res = self.tokenizer.token_decode(generate)?; - self.qwen2_5_vl.clear_kv_cache(); + self.model.clear_kv_cache(); let response = build_completion_response_with_time( res, &self.model_name, @@ -190,7 +190,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { let mut tool_call_content = String::new(); for _ in 0..sample_len { let i_start = Instant::now(); - let logits = self.qwen2_5_vl.forward( + let logits = self.model.forward( &input_ids, pixel_values, image_grid_thw, @@ -293,7 +293,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { pixel_values = None; pixel_values_video = None; } - self.qwen2_5_vl.clear_kv_cache(); + self.model.clear_kv_cache(); }; Ok(Box::new(Box::pin(stream))) } diff --git a/src/models/qwen3/generate.rs b/src/models/qwen3/generate.rs index b0c13e9..8fd3130 100644 --- a/src/models/qwen3/generate.rs +++ b/src/models/qwen3/generate.rs @@ -1,24 +1,18 @@ -use crate::models::common::MultiModalData; -use crate::models::common::generate::{ - GenerationContext, generate_generic, generate_stream_generic, -}; -use crate::params::chat::{ - ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, -}; +use crate::models::common::generate::{GenerationDataProvider, PrepareData}; + use anyhow::Result; use candle_core::{DType, Device}; use candle_nn::VarBuilder; -use rocket::futures::Stream; use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig}; use crate::models::qwen3::model::Qwen3Model; use crate::utils::{find_type_files, get_device, get_dtype}; -use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel}; +use crate::{chat_template::ChatTemplate, tokenizer::TokenizerModel}; pub struct Qwen3GenerateModel<'a> { chat_template: ChatTemplate<'a>, tokenizer: TokenizerModel, - qwen3: Qwen3Model, + model: Qwen3Model, device: Device, generation_config: Qwen3GenerationConfig, model_name: String, @@ -38,7 +32,7 @@ impl<'a> Qwen3GenerateModel<'a> { let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config: Qwen3GenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; - let qwen3 = Qwen3Model::new(&cfg, vb, generation_config.eos_token_id.clone())?; + let model = Qwen3Model::new(&cfg, vb, generation_config.eos_token_id.clone())?; let model_name = std::path::Path::new(path) .file_name() @@ -48,7 +42,7 @@ impl<'a> Qwen3GenerateModel<'a> { Ok(Qwen3GenerateModel { chat_template, tokenizer, - qwen3, + model, device: device.clone(), generation_config, model_name, @@ -56,77 +50,30 @@ impl<'a> Qwen3GenerateModel<'a> { } } -impl<'a> GenerateModel for Qwen3GenerateModel<'a> { - fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let temperature = mes - .temperature - .unwrap_or(self.generation_config.temperature); - let top_p = mes.top_p.unwrap_or(self.generation_config.top_p); - let top_k = self.generation_config.top_k; - let seed = mes.seed.unwrap_or(34562) as u64; - let mes_render = self.chat_template.apply_chat_template(&mes)?; - let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let sample_len = mes.max_tokens.unwrap_or(2048); - let mut ctx = GenerationContext::new( - temperature.into(), - top_p.into(), - top_k.into(), - mes.repeat_penalty, - mes.repeat_last_n, - seed, - input_ids.dim(1)?, - sample_len, - self.device.clone(), - ); - - let data = MultiModalData::new(vec![]); - generate_generic( - &mut self.qwen3, - &self.tokenizer, - input_ids, - data, - &mut ctx, - &self.model_name, - ) +impl<'a> GenerationDataProvider for Qwen3GenerateModel<'a> { + fn get_temperature(&self, req_temp: Option) -> Option { + Some(req_temp.unwrap_or(self.generation_config.temperature)) } - fn generate_stream( - &mut self, - mes: ChatCompletionParameters, - ) -> Result< - Box< - dyn Stream> - + Send - + Unpin - + '_, - >, - > { - let temperature = mes - .temperature - .unwrap_or(self.generation_config.temperature); - let top_p = mes.top_p.unwrap_or(self.generation_config.top_p); - let top_k = self.generation_config.top_k; - let seed = mes.seed.unwrap_or(34562) as u64; - let mes_render = self.chat_template.apply_chat_template(&mes)?; - let in_reasoning = mes_render.ends_with("\n"); + + fn get_top_p(&self, req_top_p: Option) -> Option { + Some(req_top_p.unwrap_or(self.generation_config.top_p)) + } + + fn get_top_k(&self, top_k: Option) -> Option { + Some(top_k.unwrap_or(self.generation_config.top_k)) + } + + fn get_data(&self, mes: &crate::params::chat::ChatCompletionParameters) -> Result { + let mes_render = self.chat_template.apply_chat_template(mes)?; + let in_reasoning = self.is_in_reasoning(&mes_render); let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let data = MultiModalData::new(vec![]); - let sample_len = mes.max_tokens.unwrap_or(512); - let stream = generate_stream_generic( - &mut self.qwen3, - &self.tokenizer, - input_ids, - data, - temperature.into(), - top_p.into(), - top_k.into(), - mes.repeat_penalty, - mes.repeat_last_n, - seed, - sample_len, + let multi_model_data = self.get_multi_model_data(); + Ok(PrepareData { in_reasoning, - &self.device, - &self.model_name, - )?; - Ok(Box::new(Box::pin(stream))) + input_ids, + multi_model_data, + }) } } + +crate::impl_generate_model!(Qwen3GenerateModel<'a>); diff --git a/src/models/qwen3_5/generate.rs b/src/models/qwen3_5/generate.rs index 7b56354..de97298 100644 --- a/src/models/qwen3_5/generate.rs +++ b/src/models/qwen3_5/generate.rs @@ -29,7 +29,7 @@ pub struct Qwen3_5GenerateModel<'a> { chat_template: ChatTemplate<'a>, tokenizer: TokenizerModel, pre_processor: Option, - qwen3_5: Qwen3_5Model, + model: Qwen3_5Model, device: Device, model_name: String, repeat_penalty: f32, @@ -53,13 +53,13 @@ impl<'a> Qwen3_5GenerateModel<'a> { let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; let eos_ids = vec![cfg.text_config.eos_token_id]; - let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg, eos_ids)?; + let model = Qwen3_5Model::new_from_vb(vb, cfg, eos_ids)?; Ok(Self { chat_template, tokenizer, pre_processor: Some(pre_processor), - qwen3_5, + model, device, model_name: model_name.to_string(), repeat_penalty: 1.0, @@ -89,13 +89,13 @@ impl<'a> Qwen3_5GenerateModel<'a> { let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; let eos_ids = vec![cfg.text_config.eos_token_id]; // let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg, eos_ids)?; - let qwen3_5 = Qwen3_5Model::new_from_vb_without_visual(vb, cfg, eos_ids)?; + let model = Qwen3_5Model::new_from_vb_without_visual(vb, cfg, eos_ids)?; Ok(Self { chat_template, tokenizer, pre_processor, - qwen3_5, + model, device, model_name: model_name.to_string(), repeat_penalty: 1.0, @@ -142,7 +142,7 @@ impl<'a> Qwen3_5GenerateModel<'a> { .get_matedata("tokenizer.ggml.eos_token_id")? .to_u32()?; let eos_ids = vec![eos_token_id]; - let qwen3_5 = + let model = Qwen3_5Model::new_from_gguf(&mut model_gguf, mmproj_gguf.as_mut(), &device, eos_ids)?; let stem = std::path::Path::new(model_file) .file_stem() // 获取文件名主干(不含扩展名) @@ -152,7 +152,7 @@ impl<'a> Qwen3_5GenerateModel<'a> { chat_template, tokenizer, pre_processor, - qwen3_5, + model, device, model_name: stem.to_string(), repeat_penalty: 1.2, @@ -199,13 +199,7 @@ impl<'a> Qwen3_5GenerateModel<'a> { ]; let data = MultiModalData::new(data_vec); - generate_generic_text( - &mut self.qwen3_5, - &self.tokenizer, - input_ids, - data, - &mut ctx, - ) + generate_generic_text(&mut self.model, &self.tokenizer, input_ids, data, &mut ctx) } pub fn generate_stream_text( @@ -237,7 +231,7 @@ impl<'a> Qwen3_5GenerateModel<'a> { let data = MultiModalData::new(data_vec); let seed = mes.seed.unwrap_or(34562) as u64; generate_stream_generic_text( - &mut self.qwen3_5, + &mut self.model, &self.tokenizer, input_ids, data, @@ -293,7 +287,7 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> { ]; let data = MultiModalData::new(data_vec); generate_generic( - &mut self.qwen3_5, + &mut self.model, &self.tokenizer, input_ids, data, @@ -339,7 +333,7 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> { let data = MultiModalData::new(data_vec); let seed = mes.seed.unwrap_or(34562) as u64; let stream = generate_stream_generic( - &mut self.qwen3_5, + &mut self.model, &self.tokenizer, input_ids, data, diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs index e6a5744..a01f8d4 100644 --- a/src/models/qwen3_asr/generate.rs +++ b/src/models/qwen3_asr/generate.rs @@ -37,7 +37,7 @@ pub struct Qwen3AsrGenerateModel<'a> { chat_template: ChatTemplate<'a>, tokenizer: TokenizerModel, processor: Qwen3AsrProcessor, - qwen3_asr: Qwen3ASRModel, + model: Qwen3ASRModel, device: Device, dtype: DType, eos_token_id1: u32, @@ -65,7 +65,7 @@ impl<'a> Qwen3AsrGenerateModel<'a> { let dtype = get_dtype(dtype, cfg_dtype); let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; - let qwen3_asr = Qwen3ASRModel::new(vb, &cfg, generation_config.eos_token_id.clone())?; + let model = Qwen3ASRModel::new(vb, &cfg, generation_config.eos_token_id.clone())?; let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -75,7 +75,7 @@ impl<'a> Qwen3AsrGenerateModel<'a> { chat_template, tokenizer, processor, - qwen3_asr, + model, device, dtype, eos_token_id1: generation_config.eos_token_id[0] as u32, @@ -116,13 +116,8 @@ impl<'a> Qwen3AsrGenerateModel<'a> { ); let data_vec = vec![input_features]; let data = MultiModalData::new(data_vec); - let mut text = generate_generic_text( - &mut self.qwen3_asr, - &self.tokenizer, - input_ids, - data, - &mut ctx, - )?; + let mut text = + generate_generic_text(&mut self.model, &self.tokenizer, input_ids, data, &mut ctx)?; if text.contains("") { let mut split: Vec<&str> = text.split("").collect(); text = split.pop().unwrap_or(&text).to_string(); @@ -156,7 +151,7 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> { for _ in 0..sample_len { let i_start = Instant::now(); let logits = - self.qwen3_asr + self.model .forward(&input_ids, seqlen_offset, input_features.as_ref())?; let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; let next_token = logit_processor.sample(&logits)?; @@ -175,7 +170,7 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> { input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; input_features = None; } - self.qwen3_asr.clear_kv_cache(); + self.model.clear_kv_cache(); } let num_token = generate.len() as u32; let res = self.tokenizer.token_decode(generate)?; @@ -226,7 +221,7 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> { for _ in 0..sample_len { let i_start = Instant::now(); let logits = - self.qwen3_asr + self.model .forward(&input_ids, seqlen_offset, input_features.as_ref())?; let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; let next_token = logit_processor.sample(&logits)?; @@ -266,7 +261,7 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> { input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; input_features = None; } - self.qwen3_asr.clear_kv_cache(); + self.model.clear_kv_cache(); } }; Ok(Box::new(Box::pin(stream))) diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs index b87dc6a..542fbdd 100644 --- a/src/models/qwen3vl/generate.rs +++ b/src/models/qwen3vl/generate.rs @@ -25,7 +25,7 @@ pub struct Qwen3VLGenerateModel<'a> { chat_template: ChatTemplate<'a>, tokenizer: TokenizerModel, pre_processor: Qwen3VLProcessor, - qwen3_vl: Qwen3VLModel, + model: Qwen3VLModel, device: Device, generation_config: Qwen3GenerationConfig, model_name: String, @@ -46,7 +46,7 @@ impl<'a> Qwen3VLGenerateModel<'a> { let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config: Qwen3GenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; - let qwen3_vl = Qwen3VLModel::new(cfg, vb, generation_config.eos_token_id.clone())?; + let model = Qwen3VLModel::new(cfg, vb, generation_config.eos_token_id.clone())?; let model_name = std::path::Path::new(path) .file_name() @@ -57,7 +57,7 @@ impl<'a> Qwen3VLGenerateModel<'a> { chat_template, tokenizer, pre_processor, - qwen3_vl, + model, device, generation_config, model_name, @@ -101,7 +101,7 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { ]; let data = MultiModalData::new(data_vec); generate_generic( - &mut self.qwen3_vl, + &mut self.model, &self.tokenizer, input_ids, data, @@ -145,7 +145,7 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { let data = MultiModalData::new(data_vec); let seed = mes.seed.unwrap_or(34562) as u64; let stream = generate_stream_generic( - &mut self.qwen3_vl, + &mut self.model, &self.tokenizer, input_ids, data, diff --git a/src/models/voxcpm/generate.rs b/src/models/voxcpm/generate.rs index e9ebbfb..e73bd5e 100644 --- a/src/models/voxcpm/generate.rs +++ b/src/models/voxcpm/generate.rs @@ -27,7 +27,7 @@ use crate::{ }; pub struct VoxCPMGenerate { - voxcpm: VoxCPMModel, + model: VoxCPMModel, prompt_cache: Option>, out_sample_rate: usize, model_name: String, @@ -106,12 +106,12 @@ impl VoxCPMGenerate { VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device) }; let tokenizer = SingleChineseTokenizer::new(path)?; - let voxcpm = VoxCPMModel::new(vb_voxcpm, config, tokenizer, audio_vae)?; + let model = VoxCPMModel::new(vb_voxcpm, config, tokenizer, audio_vae)?; let out_sample_rate = audio_config .out_sample_rate .unwrap_or(audio_config.sample_rate); Ok(Self { - voxcpm, + model, prompt_cache: None, out_sample_rate, model_name, @@ -124,7 +124,7 @@ impl VoxCPMGenerate { prompt_wav_path: String, ) -> Result<()> { let cache = self - .voxcpm + .model .build_prompt_cache(prompt_text, prompt_wav_path)?; self.prompt_cache = Some(cache); Ok(()) @@ -143,7 +143,7 @@ impl VoxCPMGenerate { let audio = match &self.prompt_cache { Some(cache) => { let prompt_cache = cache.clone(); - self.voxcpm.generate_with_prompt_cache( + self.model.generate_with_prompt_cache( target_text, prompt_cache, min_len, @@ -156,7 +156,7 @@ impl VoxCPMGenerate { } None => self.generate_simple(target_text)?, }; - self.voxcpm.clear_kv_cache(); + self.model.clear_kv_cache(); Ok(audio) } @@ -196,7 +196,7 @@ impl VoxCPMGenerate { // retry_badcase: bool, retry_badcase_ratio_threshold: f64, ) -> Result { - let audio = self.voxcpm.generate( + let audio = self.model.generate( target_text, prompt_text, prompt_wav_path, @@ -207,7 +207,7 @@ impl VoxCPMGenerate { // retry_badcase, retry_badcase_ratio_threshold, )?; - self.voxcpm.clear_kv_cache(); + self.model.clear_kv_cache(); Ok(audio) } @@ -250,7 +250,7 @@ impl GenerateModel for VoxCPMGenerate { target_text = format!("({instruction}){target_text}"); } let audio = self - .voxcpm + .model .generate( target_text, prompt_text, @@ -262,13 +262,13 @@ impl GenerateModel for VoxCPMGenerate { retry_badcase_ratio_threshold, ) .inspect_err(|_| { - self.voxcpm.clear_kv_cache(); + self.model.clear_kv_cache(); })?; let wav_u8 = get_audio_wav_u8(&audio, self.out_sample_rate as u32)?; // let wave_u8_str = String::from_utf8(wav_u8)?; let base64_audio = BASE64_STANDARD.encode(wav_u8); let response = build_audio_completion_response(&base64_audio, &self.model_name); - self.voxcpm.clear_kv_cache(); + self.model.clear_kv_cache(); Ok(response) } #[allow(unused_variables)] diff --git a/src/models/voxcpm_refact/generate.rs b/src/models/voxcpm_refact/generate.rs index c5bd0f1..9ae2419 100644 --- a/src/models/voxcpm_refact/generate.rs +++ b/src/models/voxcpm_refact/generate.rs @@ -24,7 +24,7 @@ use crate::{ }; pub struct VoxCPMGenerateRefact { - voxcpm: VoxCPMModelRefact, + model: VoxCPMModelRefact, tokenizer: SingleChineseTokenizer, audio_vae: AudioVAE, processor: VoxCPMProcessor, @@ -113,13 +113,13 @@ impl VoxCPMGenerateRefact { VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device) }; let tokenizer = SingleChineseTokenizer::new(path)?; - let voxcpm = + let model = VoxCPMModelRefact::new(vb_voxcpm, config, audio_vae.latent_dim, decode_chunk_size)?; let out_sample_rate = audio_config .out_sample_rate .unwrap_or(audio_config.sample_rate); Ok(Self { - voxcpm, + model, tokenizer, audio_vae, processor, @@ -162,7 +162,7 @@ impl VoxCPMGenerateRefact { } else { max_len }; - let audio = self.voxcpm.inference( + let audio = self.model.inference( &text_token, audio_feat.as_ref(), audio_mask.as_ref(), @@ -172,7 +172,7 @@ impl VoxCPMGenerateRefact { cfg_value, &self.audio_vae, )?; - self.voxcpm.clear_kv_cache(); + self.model.clear_kv_cache(); Ok(audio) } @@ -240,7 +240,7 @@ impl VoxCPMGenerateRefact { } else { max_len }; - self.voxcpm.inference( + self.model.inference( &text_token, audio_feat.as_ref(), audio_mask.as_ref(), @@ -255,7 +255,7 @@ impl VoxCPMGenerateRefact { return Err(anyhow!("need prompt_cache")); } }; - self.voxcpm.clear_kv_cache(); + self.model.clear_kv_cache(); Ok(audio) } @@ -284,7 +284,7 @@ impl VoxCPMGenerateRefact { } else { max_len }; - self.voxcpm.inference_stream( + self.model.inference_stream( text_token, audio_feat, audio_mask, @@ -344,13 +344,13 @@ impl GenerateModel for VoxCPMGenerateRefact { retry_badcase_ratio_threshold, ) .inspect_err(|_| { - self.voxcpm.clear_kv_cache(); + self.model.clear_kv_cache(); })?; let wav_u8 = get_audio_wav_u8(&audio, self.out_sample_rate as u32)?; // let wave_u8_str = String::from_utf8(wav_u8)?; let base64_audio = BASE64_STANDARD.encode(wav_u8); let response = build_audio_completion_response(&base64_audio, &self.model_name); - self.voxcpm.clear_kv_cache(); + self.model.clear_kv_cache(); Ok(response) } #[allow(unused_variables)] diff --git a/tests/test_minicpm5.rs b/tests/test_minicpm5.rs index 9cea524..f429701 100644 --- a/tests/test_minicpm5.rs +++ b/tests/test_minicpm5.rs @@ -1,10 +1,11 @@ -use std::time::Instant; +use std::{pin::pin, time::Instant}; use aha::{ models::{GenerateModel, minicpm5::generate::MiniCPM5GenerateModel}, params::chat::ChatCompletionParameters, }; use anyhow::Result; +use rocket::futures::StreamExt; #[test] fn minicpm5_generate() -> Result<()> { // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_minicpm5 minicpm5_generate -r -- --nocapture @@ -16,7 +17,7 @@ fn minicpm5_generate() -> Result<()> { { "temperature": 0.3, "top_p": 0.8, - "model": "minicpm4", + "model": "minicpm5", "messages": [ { "role": "user", @@ -39,3 +40,41 @@ fn minicpm5_generate() -> Result<()> { } Ok(()) } + +#[tokio::test] +async fn minicpm5_stream() -> Result<()> { + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_minicpm5 minicpm5_stream -r -- --nocapture + + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/OpenBMB/MiniCPM5-1B/", save_dir); + + let message = r#" + { + "model": "minicpm5", + "messages": [ + { + "role": "user", + "content": "什么是AI" + } + ], + "enable_thinking": true + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut model = MiniCPM5GenerateModel::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 mut stream = pin!(model.generate_stream(mes)?); + while let Some(item) = stream.next().await { + println!("generate: \n {:?}", item); + } + + let i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + + Ok(()) +} diff --git a/tests/test_qwen3.rs b/tests/test_qwen3.rs index 895daf0..26ae362 100644 --- a/tests/test_qwen3.rs +++ b/tests/test_qwen3.rs @@ -42,7 +42,7 @@ fn qwen3_0_6b_generate() -> Result<()> { #[tokio::test] async fn qwen3_0_6b_stream() -> Result<()> { - // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda qwen3_0_6b_stream -r -- --nocapture + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_qwen3 qwen3_0_6b_stream -r -- --nocapture let save_dir = aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;