diff --git a/README.md b/README.md index b0401c8..d3def6c 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-29 +- generate code refactored + ### 2026-05-28 - generate code refactoring progress 1/3 @@ -54,10 +57,6 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an ### 2026-05-09 - merge pr/eastgold15/46, add aha-ui -### 2026-04-25 -- VoxCPM update stream - - **[View full changelog](docs/changelog.md)** → ## Why aha? diff --git a/README.zh-CN.md b/README.zh-CN.md index 7dd64ee..0b46f03 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -38,6 +38,9 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理 | **重排序** | Qwen3-Reranker | ## 更新日志 +### 2026-05-29 +- generate代码重构完成 + ### 2026-05-28 - generate代码重构进度 1/3 @@ -53,9 +56,6 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理 ### 2026-05-09 - 合并 pr/eastgold15/46, 添加 aha-ui -### 2026-04-25 -- VoxCPM 更新流式生成 - **[查看完整更新日志](docs/changelog.zh-CN.md)** → diff --git a/docs/changelog.md b/docs/changelog.md index 38b07b0..7450b46 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-29 +- generate code refactored + ### 2026-05-28 - generate code refactoring progress 1/3 diff --git a/docs/changelog.zh-CN.md b/docs/changelog.zh-CN.md index a4b0279..4809833 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-29 +- generate代码重构完成 + ### 2026-05-28 - generate代码重构进度 1/3 diff --git a/src/models/fun_asr_nano/generate.rs b/src/models/fun_asr_nano/generate.rs index 75a8cc2..9064704 100644 --- a/src/models/fun_asr_nano/generate.rs +++ b/src/models/fun_asr_nano/generate.rs @@ -3,18 +3,16 @@ use std::collections::HashMap; use crate::{ models::common::{ MultiModalData, - generate::{GenerationContext, generate_generic, generate_stream_generic}, + generate::{GenerationDataProvider, PrepareData}, }, - params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, + params::chat::ChatCompletionParameters, }; use anyhow::{Result, anyhow}; use candle_core::{DType, Device, pickle::read_all_with_key}; use candle_nn::VarBuilder; -use rocket::futures::Stream; use crate::{ models::{ - GenerateModel, fun_asr_nano::{ config::FunASRNanoConfig, model::FunAsrNanoModel, processor::FunAsrNanoProcessor, }, @@ -94,79 +92,30 @@ impl FunAsrNanoGenerateModel { } } -impl GenerateModel for FunAsrNanoGenerateModel { - 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 max_tokens = mes.max_tokens.unwrap_or(1024); - let (speech, fbank_mask, input_ids) = self.processor.process_info(&mes, &self.tokenizer)?; - let speech = speech.to_dtype(self.dtype)?; - 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)?, - max_tokens, - self.device.clone(), - ); - - let data_vec = vec![speech.into(), fbank_mask.into()]; - let data = MultiModalData::new(data_vec); - generate_generic( - &mut self.model, - &self.tokenizer, - input_ids, - data, - &mut ctx, - &self.model_name, - ) +impl GenerationDataProvider for FunAsrNanoGenerateModel { + 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 max_tokens = mes.max_tokens.unwrap_or(1024); - let (speech, fbank_mask, input_ids) = self.processor.process_info(&mes, &self.tokenizer)?; + 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: &ChatCompletionParameters) -> Result { + let (speech, fbank_mask, input_ids) = self.processor.process_info(mes, &self.tokenizer)?; let speech = speech.to_dtype(self.dtype)?; let data_vec = vec![speech.into(), fbank_mask.into()]; - let data = MultiModalData::new(data_vec); - let stream = generate_stream_generic( - &mut self.model, - &self.tokenizer, + let multi_model_data = MultiModalData::new(data_vec); + + Ok(PrepareData { + in_reasoning: false, input_ids, - data, - temperature.into(), - top_p.into(), - top_k.into(), - 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!(FunAsrNanoGenerateModel); diff --git a/src/models/glm_asr_nano/generate.rs b/src/models/glm_asr_nano/generate.rs index 81d3a17..77cbf64 100644 --- a/src/models/glm_asr_nano/generate.rs +++ b/src/models/glm_asr_nano/generate.rs @@ -1,22 +1,18 @@ use crate::{ models::common::{ MultiModalData, - generate::{GenerationContext, generate_generic, generate_stream_generic}, + generate::{GenerationDataProvider, PrepareData}, }, - params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, + params::chat::ChatCompletionParameters, }; use anyhow::Result; use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; -use rocket::futures::Stream; use crate::{ chat_template::ChatTemplate, - models::{ - GenerateModel, - glm_asr_nano::{ - config::GlmAsrNanoConfig, model::GlmAsrNanoModel, processor::GlmAsrNanoProcessor, - }, + models::glm_asr_nano::{ + config::GlmAsrNanoConfig, model::GlmAsrNanoModel, processor::GlmAsrNanoProcessor, }, tokenizer::TokenizerModel, utils::{find_type_files, get_device, get_dtype}, @@ -63,77 +59,22 @@ impl<'a> GlmAsrNanoGenerateModel<'a> { } } -impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> { - fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let seed = mes.seed.unwrap_or(34562) as u64; - let render_text: String = self.chat_template.apply_chat_template(&mes)?; +impl<'a> GenerationDataProvider for GlmAsrNanoGenerateModel<'a> { + fn get_data(&self, mes: &ChatCompletionParameters) -> Result { + let render_text: String = self.chat_template.apply_chat_template(mes)?; let (input_features, audio_token_lengths, replace_text) = - self.processor.process_info(&mes, &render_text)?; + self.processor.process_info(mes, &render_text)?; let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; let input_features = input_features.to_dtype(self.dtype)?; let audio_token_lengths = Tensor::new(audio_token_lengths, &self.device)?; - let sample_len = mes.max_tokens.unwrap_or(1024); - let mut ctx = GenerationContext::new( - mes.temperature, - mes.top_p, - mes.top_k, - mes.repeat_penalty, - mes.repeat_last_n, - seed, - input_ids.dim(1)?, - sample_len, - self.device.clone(), - ); - let data_vec = vec![input_features.into(), audio_token_lengths.into()]; - let data = MultiModalData::new(data_vec); - generate_generic( - &mut self.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 seed = mes.seed.unwrap_or(34562) as u64; - let render_text = self.chat_template.apply_chat_template(&mes)?; - let (input_features, audio_token_lengths, replace_text) = - self.processor.process_info(&mes, &render_text)?; - let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; - let input_features = input_features.to_dtype(self.dtype)?; - let audio_token_lengths = Tensor::new(audio_token_lengths, &self.device)?; - let sample_len = mes.max_tokens.unwrap_or(1024); - let data_vec = vec![input_features.into(), audio_token_lengths.into()]; - let data = MultiModalData::new(data_vec); - 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!(GlmAsrNanoGenerateModel<'a>); diff --git a/src/models/glm_ocr/generate.rs b/src/models/glm_ocr/generate.rs index 1b3b6df..3ea262e 100644 --- a/src/models/glm_ocr/generate.rs +++ b/src/models/glm_ocr/generate.rs @@ -2,33 +2,26 @@ use crate::{ models::common::{ MultiModalData, - generate::{GenerationContext, generate_generic, generate_stream_generic}, + generate::{GenerationDataProvider, PrepareData}, }, - params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, + params::chat::ChatCompletionParameters, }; -use anyhow::{Result, anyhow}; -use candle_core::{DType, Device}; -use candle_nn::VarBuilder; -use rocket::futures::Stream; - use crate::{ - // chat_template::ChatTemplate, - models::{ - GenerateModel, - glm_ocr::{ - config::{GlmOcrConfig, GlmOcrGenerationConfig}, - model::GlmOcrModel, - processor::GlmOcrProcessor, - }, + models::glm_ocr::{ + config::{GlmOcrConfig, GlmOcrGenerationConfig}, + model::GlmOcrModel, + processor::GlmOcrProcessor, }, tokenizer::TokenizerModel, utils::{ extract_user_text, find_type_files, get_device, get_dtype, img_utils::extract_image_url, }, }; +use anyhow::{Result, anyhow}; +use candle_core::{DType, Device}; +use candle_nn::VarBuilder; pub struct GlmOcrGenerateModel { - // chat_template: ChatTemplate<'a>, tokenizer: TokenizerModel, processor: GlmOcrProcessor, model: GlmOcrModel, @@ -44,7 +37,6 @@ pub struct GlmOcrGenerateModel { impl GlmOcrGenerateModel { pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { - // let chat_template = ChatTemplate::init(path)?; let tokenizer = TokenizerModel::init(path)?; let config_path = path.to_string() + "/config.json"; let cfg: GlmOcrConfig = serde_json::from_slice(&std::fs::read(config_path)?)?; @@ -64,7 +56,6 @@ impl GlmOcrGenerateModel { .unwrap_or("ZhipuAI/GLM-OCR") .to_string(); Ok(Self { - // chat_template, tokenizer, processor, model, @@ -80,17 +71,15 @@ impl GlmOcrGenerateModel { } } -impl GenerateModel for GlmOcrGenerateModel { - fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let seed = mes.seed.unwrap_or(34562) as u64; - // Extract image path and prompt from messages - let image_urls = extract_image_url(&mes); +impl GenerationDataProvider for GlmOcrGenerateModel { + fn get_data(&self, mes: &ChatCompletionParameters) -> Result { + let image_urls = extract_image_url(mes); let image_path = image_urls .first() .ok_or_else(|| anyhow!("No image provided"))?; // Get prompt text from messages - let mut prompt = extract_user_text(&mes)?; + let mut prompt = extract_user_text(mes)?; if prompt.chars().count() == 0 { prompt = "Extract all text from this image.".to_string() } @@ -108,95 +97,18 @@ impl GenerateModel for GlmOcrGenerateModel { )?; let input_ids = processed.input_ids; - let sample_len = mes.max_tokens.unwrap_or(1024); - let mut ctx = GenerationContext::new( - mes.temperature, - mes.top_p, - mes.top_k, - mes.repeat_penalty, - mes.repeat_last_n, - seed, - input_ids.dim(1)?, - sample_len, - self.device.clone(), - ); - let data_vec = vec![ processed.pixel_values.into(), processed.grid_thw.into(), processed.image_mask.into(), ]; - let data = MultiModalData::new(data_vec); - generate_generic( - &mut self.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 seed = mes.seed.unwrap_or(34562) as u64; - // Extract image path and prompt from messages - let image_urls = extract_image_url(&mes); - let image_path = image_urls - .first() - .ok_or_else(|| anyhow!("No image provided"))?; - - // Get prompt text from messages - let mut prompt = extract_user_text(&mes)?; - if prompt.chars().count() == 0 { - prompt = "Extract all text from this image.".to_string() - } - - let processed = self.processor.process_info( - image_path, - &prompt, - &self.tokenizer, - self.image_token_id, - self.image_start_token_id, - self.image_end_token_id, - self.patch_size, - self.temporal_patch_size, - self.spatial_merge_size, - )?; - - let input_ids = processed.input_ids; - let sample_len = mes.max_tokens.unwrap_or(1024); - let data_vec = vec![ - processed.pixel_values.into(), - processed.grid_thw.into(), - processed.image_mask.into(), - ]; - let data = MultiModalData::new(data_vec); - 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!(GlmOcrGenerateModel); diff --git a/src/models/hunyuan_ocr/generate.rs b/src/models/hunyuan_ocr/generate.rs index 7e49f84..13a411f 100644 --- a/src/models/hunyuan_ocr/generate.rs +++ b/src/models/hunyuan_ocr/generate.rs @@ -1,24 +1,20 @@ use crate::{ models::common::{ MultiModalData, - generate::{GenerationContext, generate_generic, generate_stream_generic}, + generate::{GenerationDataProvider, PrepareData}, }, - params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, + params::chat::ChatCompletionParameters, }; use anyhow::Result; use candle_core::{DType, Device}; use candle_nn::VarBuilder; -use rocket::futures::Stream; use crate::{ chat_template::ChatTemplate, - models::{ - GenerateModel, - hunyuan_ocr::{ - config::{HunYuanVLConfig, HunyuanOCRGenerationConfig}, - model::HunyuanVLModel, - processor::HunyuanVLProcessor, - }, + models::hunyuan_ocr::{ + config::{HunYuanVLConfig, HunyuanOCRGenerationConfig}, + model::HunyuanVLModel, + processor::HunyuanVLProcessor, }, tokenizer::TokenizerModel, utils::{find_type_files, get_device, get_dtype}, @@ -68,72 +64,24 @@ impl<'a> HunyuanOCRGenerateModel<'a> { } } -impl<'a> GenerateModel for HunyuanOCRGenerateModel<'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 data = self - .pre_processor - .process_info(&mes, &self.tokenizer, &mes_render)?; - let sample_len = mes.max_tokens.unwrap_or(1024); - let input_ids = data.input_ids; - 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_vec = vec![ - data.pixel_values, - data.image_grid_thw, - data.image_mask.into(), - data.position_ids.into(), - ]; - let data = MultiModalData::new(data_vec); - generate_generic( - &mut self.model, - &self.tokenizer, - input_ids, - data, - &mut ctx, - &self.model_name, - ) +impl<'a> GenerationDataProvider for HunyuanOCRGenerateModel<'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)?; + 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: &ChatCompletionParameters) -> Result { + let mes_render = self.chat_template.apply_chat_template(mes)?; let data = self .pre_processor - .process_info(&mes, &self.tokenizer, &mes_render)?; - - let sample_len = mes.max_tokens.unwrap_or(1024); + .process_info(mes, &self.tokenizer, &mes_render)?; let input_ids = data.input_ids; let data_vec = vec![ data.pixel_values, @@ -141,23 +89,13 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> { data.image_mask.into(), data.position_ids.into(), ]; - let data = MultiModalData::new(data_vec); - let stream = generate_stream_generic( - &mut self.model, - &self.tokenizer, + let multi_model_data = MultiModalData::new(data_vec); + Ok(PrepareData { + in_reasoning: false, input_ids, - data, - temperature.into(), - top_p.into(), - top_k.into(), - 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!(HunyuanOCRGenerateModel<'a>); diff --git a/src/models/lfm2vl/generate.rs b/src/models/lfm2vl/generate.rs index afa7183..d56b48a 100644 --- a/src/models/lfm2vl/generate.rs +++ b/src/models/lfm2vl/generate.rs @@ -1,9 +1,9 @@ use crate::{ models::common::{ MultiModalData, - generate::{GenerationContext, generate_generic, generate_stream_generic}, + generate::{GenerationDataProvider, PrepareData}, }, - params::chat::{ChatCompletionParameters, ChatCompletionResponse}, + params::chat::ChatCompletionParameters, }; use anyhow::Result; use candle_core::{DType, Device}; @@ -12,7 +12,6 @@ use candle_nn::VarBuilder; use crate::{ chat_template::ChatTemplate, models::{ - GenerateModel, lfm2::config::Lfm2GenerateConfig, lfm2vl::{config::Lfm2VLConfig, model::Lfm2VLModel, processor::Lfm2VLProcessor}, }, @@ -60,82 +59,24 @@ impl<'a> Lfm2VLGenerateModel<'a> { } } -impl<'a> GenerateModel for Lfm2VLGenerateModel<'a> { - fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let mes_render = self.chat_template.apply_chat_template(&mes)?; +impl<'a> GenerationDataProvider for Lfm2VLGenerateModel<'a> { + fn get_data(&self, mes: &ChatCompletionParameters) -> Result { + let mes_render = self.chat_template.apply_chat_template(mes)?; let (pixel_values, pixel_attention_mask, spatial_shapes, text) = - self.processor.process_info(&mes, &mes_render)?; + self.processor.process_info(mes, &mes_render)?; let input_ids = self.tokenizer.text_encode(text, &self.device)?; - let seed = mes.seed.unwrap_or(34562) as u64; - let sample_len = mes.max_tokens.unwrap_or(1024); - 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_vec = vec![ pixel_values.into(), pixel_attention_mask.into(), spatial_shapes.into(), ]; - let data = MultiModalData::new(data_vec); - generate_generic( - &mut self.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 rocket::futures::Stream< - Item = Result, - > + Send - + Unpin - + '_, - >, - > { - let mes_render = self.chat_template.apply_chat_template(&mes)?; - let (pixel_values, pixel_attention_mask, spatial_shapes, text) = - self.processor.process_info(&mes, &mes_render)?; - let input_ids = self.tokenizer.text_encode(text, &self.device)?; - let sample_len = mes.max_tokens.unwrap_or(1024); - let data_vec = vec![ - pixel_values.into(), - pixel_attention_mask.into(), - spatial_shapes.into(), - ]; - let data = MultiModalData::new(data_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!(Lfm2VLGenerateModel<'a>); diff --git a/src/models/paddleocr_vl/generate.rs b/src/models/paddleocr_vl/generate.rs index a2a6a29..b39432c 100644 --- a/src/models/paddleocr_vl/generate.rs +++ b/src/models/paddleocr_vl/generate.rs @@ -1,21 +1,16 @@ 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 crate::params::chat::ChatCompletionParameters; use anyhow::Result; use candle_core::{D, DType, Device, IndexOp, Tensor}; use candle_nn::VarBuilder; -use rocket::futures::Stream; use crate::models::paddleocr_vl::config::{PaddleOCRVLConfig, PaddleOCRVLPreprocessorConfig}; use crate::models::paddleocr_vl::model::PaddleOCRVLModel; use crate::models::paddleocr_vl::processor::PaddleOCRVLProcessor; use crate::utils::tensor_utils::get_equal_mask; 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 PaddleOCRVLGenerateModel<'a> { chat_template: ChatTemplate<'a>, @@ -60,11 +55,11 @@ impl<'a> PaddleOCRVLGenerateModel<'a> { } } -impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> { - fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let mes_render = self.chat_template.apply_chat_template(&mes)?; +impl<'a> GenerationDataProvider for PaddleOCRVLGenerateModel<'a> { + fn get_data(&self, mes: &ChatCompletionParameters) -> Result { + let mes_render = self.chat_template.apply_chat_template(mes)?; let (replace_text, pixel_values, image_grid_thw) = - self.pre_processor.process_info(&mes, &mes_render)?; + self.pre_processor.process_info(mes, &mes_render)?; let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; let image_mask = get_equal_mask(&input_ids, self.cfg.image_token_id)?; @@ -73,84 +68,19 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> { .cumsum(D::Minus1)? .to_dtype(candle_core::DType::U32)? .broadcast_sub(&Tensor::new(vec![1_u32], input_ids.device())?)?; - let seed = mes.seed.unwrap_or(34562) as u64; - let sample_len = mes.max_tokens.unwrap_or(1024); - 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_vec = vec![ pixel_values, image_grid_thw, image_mask.into(), cache_position.into(), ]; - let data = MultiModalData::new(data_vec); - generate_generic( - &mut self.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 mes_render = self.chat_template.apply_chat_template(&mes)?; - let (replace_text, pixel_values, image_grid_thw) = - self.pre_processor.process_info(&mes, &mes_render)?; - let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; - let image_mask = get_equal_mask(&input_ids, self.cfg.image_token_id)?; - - let cache_position = Tensor::ones_like(&input_ids.i(0)?)? - .to_dtype(candle_core::DType::F64)? - .cumsum(D::Minus1)? - .to_dtype(candle_core::DType::U32)? - .broadcast_sub(&Tensor::new(vec![1_u32], input_ids.device())?)?; - - let sample_len = mes.max_tokens.unwrap_or(1024); - let data_vec = vec![ - pixel_values, - image_grid_thw, - image_mask.into(), - cache_position.into(), - ]; - let data = MultiModalData::new(data_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!(PaddleOCRVLGenerateModel<'a>); diff --git a/src/models/qwen3_5/generate.rs b/src/models/qwen3_5/generate.rs index de97298..7028f53 100644 --- a/src/models/qwen3_5/generate.rs +++ b/src/models/qwen3_5/generate.rs @@ -2,11 +2,11 @@ use crate::{ models::common::{ MultiModalData, generate::{ - GenerationContext, generate_generic, generate_generic_text, generate_stream_generic, + GenerationContext, GenerationDataProvider, PrepareData, generate_generic_text, generate_stream_generic_text, }, }, - params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, + params::chat::ChatCompletionParameters, }; use anyhow::{Result, anyhow}; use candle_core::{DType, Device, quantized::gguf_file}; @@ -16,7 +16,6 @@ use rocket::futures::Stream; use crate::{ chat_template::ChatTemplate, models::{ - GenerateModel, common::gguf::Gguf, qwen3_5::{config::Qwen3_5Config, model::Qwen3_5Model}, qwen3vl::processor::Qwen3VLProcessor, @@ -247,71 +246,25 @@ impl<'a> Qwen3_5GenerateModel<'a> { } } -impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> { - fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let seed = mes.seed.unwrap_or(32768) as u64; - let temperature = mes.temperature.unwrap_or(0.4); - let top_p = mes.top_p.unwrap_or(0.95); - let mes_render = self.chat_template.apply_chat_template(&mes)?; - let (mes_text, pixel_values, image_grid_thw, pixel_values_video, video_grid_thw) = - if let Some(processor) = &self.pre_processor { - let input = processor.process_info(&mes, &mes_render)?; - ( - input.replace_text, - input.pixel_values, - input.image_grid_thw, - input.pixel_values_video, - input.video_grid_thw, - ) - } else { - (mes_render, None, None, None, None) - }; - let input_ids = self.tokenizer.text_encode(mes_text, &self.device)?; - let sample_len = mes.max_tokens.unwrap_or(1024); - let mut ctx = GenerationContext::new( - temperature.into(), - top_p.into(), - Some(20), - self.repeat_penalty.into(), - self.repeat_last_n.into(), - seed, - input_ids.dim(1)?, - sample_len, - self.device.clone(), - ); - let data_vec = vec![ - pixel_values, - image_grid_thw, - pixel_values_video, - video_grid_thw, - ]; - let data = MultiModalData::new(data_vec); - generate_generic( - &mut self.model, - &self.tokenizer, - input_ids, - data, - &mut ctx, - &self.model_name, - ) +impl<'a> GenerationDataProvider for Qwen3_5GenerateModel<'a> { + fn get_temperature(&self, req_temp: Option) -> Option { + Some(req_temp.unwrap_or(0.4)) } - fn generate_stream( - &mut self, - mes: ChatCompletionParameters, - ) -> Result< - Box< - dyn Stream> - + Send - + Unpin - + '_, - >, - > { - 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(0.95)) + } + + fn get_top_k(&self, top_k: Option) -> Option { + Some(top_k.unwrap_or(40)) + } + + fn get_data(&self, mes: &ChatCompletionParameters) -> Result { + let mes_render = self.chat_template.apply_chat_template(mes)?; + let in_reasoning = self.is_in_reasoning(&mes_render); let (mes_text, pixel_values, image_grid_thw, pixel_values_video, video_grid_thw) = if let Some(processor) = &self.pre_processor { - let input = processor.process_info(&mes, &mes_render)?; + let input = processor.process_info(mes, &mes_render)?; ( input.replace_text, input.pixel_values, @@ -323,31 +276,19 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> { (mes_render, None, None, None, None) }; let input_ids = self.tokenizer.text_encode(mes_text, &self.device)?; - let sample_len = mes.max_tokens.unwrap_or(1024); let data_vec = vec![ pixel_values, image_grid_thw, pixel_values_video, video_grid_thw, ]; - let data = MultiModalData::new(data_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, - self.repeat_penalty.into(), - self.repeat_last_n.into(), - seed, - sample_len, + let multi_model_data = MultiModalData::new(data_vec); + Ok(PrepareData { in_reasoning, - &self.device, - &self.model_name, - )?; - Ok(Box::new(Box::pin(stream))) + input_ids, + multi_model_data, + }) } } + +crate::impl_generate_model!(Qwen3_5GenerateModel<'a>); diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs index 542fbdd..39caf7d 100644 --- a/src/models/qwen3vl/generate.rs +++ b/src/models/qwen3vl/generate.rs @@ -1,19 +1,17 @@ use crate::{ models::common::{ MultiModalData, - generate::{GenerationContext, generate_generic, generate_stream_generic}, + generate::{GenerationDataProvider, PrepareData}, }, - params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, + params::chat::ChatCompletionParameters, }; use anyhow::Result; use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; -use rocket::futures::Stream; use crate::{ chat_template::ChatTemplate, models::{ - GenerateModel, qwen3::config::Qwen3GenerationConfig, qwen3vl::{config::Qwen3VLConfig, model::Qwen3VLModel, processor::Qwen3VLProcessor}, }, @@ -65,76 +63,28 @@ impl<'a> Qwen3VLGenerateModel<'a> { } } -impl<'a> GenerateModel for Qwen3VLGenerateModel<'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 = self.pre_processor.process_info(&mes, &mes_render)?; - let input_ids = self - .tokenizer - .text_encode(input.replace_text.clone(), &self.device)?; - let seq_len = input_ids.dim(1)?; - let sample_len = mes.max_tokens.unwrap_or(1024); - 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 cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?; - let data_vec = vec![ - input.pixel_values, - input.image_grid_thw, - input.pixel_values_video, - input.video_grid_thw, - cache_position.into(), - ]; - let data = MultiModalData::new(data_vec); - generate_generic( - &mut self.model, - &self.tokenizer, - input_ids, - data, - &mut ctx, - &self.model_name, - ) +impl<'a> GenerationDataProvider for Qwen3VLGenerateModel<'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 mes_render = self.chat_template.apply_chat_template(&mes)?; - let in_reasoning = mes_render.ends_with("\n"); - let input = self.pre_processor.process_info(&mes, &mes_render)?; + 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: &ChatCompletionParameters) -> Result { + let mes_render = self.chat_template.apply_chat_template(mes)?; + let in_reasoning = self.is_in_reasoning(&mes_render); + let input = self.pre_processor.process_info(mes, &mes_render)?; let input_ids = self .tokenizer .text_encode(input.replace_text.clone(), &self.device)?; let seq_len = input_ids.dim(1)?; let cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?; - let sample_len = mes.max_tokens.unwrap_or(1024); let data_vec = vec![ input.pixel_values, input.image_grid_thw, @@ -142,24 +92,13 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { input.video_grid_thw, cache_position.into(), ]; - let data = MultiModalData::new(data_vec); - let seed = mes.seed.unwrap_or(34562) as u64; - let stream = generate_stream_generic( - &mut self.model, - &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 = MultiModalData::new(data_vec); + Ok(PrepareData { in_reasoning, - &self.device, - &self.model_name, - )?; - Ok(Box::new(Box::pin(stream))) + input_ids, + multi_model_data, + }) } } + +crate::impl_generate_model!(Qwen3VLGenerateModel<'a>); diff --git a/tests/test_fun_asr_nano.rs b/tests/test_fun_asr_nano.rs index cfec80b..7418d03 100644 --- a/tests/test_fun_asr_nano.rs +++ b/tests/test_fun_asr_nano.rs @@ -21,7 +21,7 @@ fn fun_asr_nano_generate() -> Result<()> { "type": "audio", "audio_url": { - "url": "file://./assets/audio/voice_01.wav" + "url": "https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav" } }, { diff --git a/tests/test_glm_asr_nano.rs b/tests/test_glm_asr_nano.rs index 110e48a..bbe520e 100644 --- a/tests/test_glm_asr_nano.rs +++ b/tests/test_glm_asr_nano.rs @@ -22,7 +22,7 @@ fn glm_asr_nano_generate() -> Result<()> { "type": "audio", "audio_url": { - "url": "file://./assets/audio/voice_01.wav" + "url": "https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav" } }, { diff --git a/tests/test_glm_ocr.rs b/tests/test_glm_ocr.rs index 78ca103..fe79993 100644 --- a/tests/test_glm_ocr.rs +++ b/tests/test_glm_ocr.rs @@ -52,7 +52,7 @@ fn glm_ocr_generate() -> Result<()> { #[tokio::test] async fn glm_ocr_stream() -> Result<()> { - // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda glm_ocr_stream -r -- --nocapture + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_glm_ocr glm_ocr_stream -r -- --nocapture let message = r#" { diff --git a/tests/test_hunyuan_ocr.rs b/tests/test_hunyuan_ocr.rs index 48242e1..f03abc5 100644 --- a/tests/test_hunyuan_ocr.rs +++ b/tests/test_hunyuan_ocr.rs @@ -49,7 +49,7 @@ fn hunyuan_ocr_generate() -> Result<()> { #[tokio::test] async fn hunyuan_ocr_stream() -> Result<()> { - // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda hunyuan_ocr_stream -r -- --nocapture + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_hunyuan_ocr hunyuan_ocr_stream -r -- --nocapture let message = r#" { diff --git a/tests/test_lfm2vl.rs b/tests/test_lfm2vl.rs index 6a3bab4..5f57e0f 100644 --- a/tests/test_lfm2vl.rs +++ b/tests/test_lfm2vl.rs @@ -60,7 +60,7 @@ async fn lfm2vl_stream() -> Result<()> { let save_dir = aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; // let model_path = format!("{}/LiquidAI/LFM2-1.2B/", save_dir); - let model_path = format!("{}/LiquidAI/LFM2.5-VL-1.6B/", save_dir); + let model_path = format!("{}/LiquidAI/LFM2.5-VL-450M/", save_dir); let message = r#" { "model": "lfm2vl", diff --git a/tests/test_qwen3vl.rs b/tests/test_qwen3vl.rs index 7d0fd61..8d6601a 100644 --- a/tests/test_qwen3vl.rs +++ b/tests/test_qwen3vl.rs @@ -104,7 +104,7 @@ fn qwen3vl_generate() -> Result<()> { #[tokio::test] async fn qwen3vl_stream() -> Result<()> { - // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda,ffmpeg qwen3vl_stream -r -- --nocapture + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_qwen3vl qwen3vl_stream -r -- --nocapture let save_dir = aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; @@ -118,15 +118,15 @@ async fn qwen3vl_stream() -> Result<()> { "role": "user", "content": [ { - "type": "video", - "video_url": + "type": "image", + "image_url": { - "url": "./assets/video/video_test.mp4" + "url": "file://./assets/img/ocr_test1.png" } - }, + }, { "type": "text", - "text": "视频中发生了什么?" + "text": "OCR" } ] }