generate code refactored

This commit is contained in:
jhqxxx
2026-05-29 22:10:40 +08:00
parent 257057de07
commit 791f3ea7e3
18 changed files with 182 additions and 686 deletions
+3 -4
View File
@@ -39,6 +39,9 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
| **Reranker** | Qwen3-Reranker | | **Reranker** | Qwen3-Reranker |
## Changelog ## Changelog
### 2026-05-29
- generate code refactored
### 2026-05-28 ### 2026-05-28
- generate code refactoring progress 1/3 - 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 ### 2026-05-09
- merge pr/eastgold15/46, add aha-ui - merge pr/eastgold15/46, add aha-ui
### 2026-04-25
- VoxCPM update stream
**[View full changelog](docs/changelog.md)** → **[View full changelog](docs/changelog.md)** →
## Why aha? ## Why aha?
+3 -3
View File
@@ -38,6 +38,9 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
| **重排序** | Qwen3-Reranker | | **重排序** | Qwen3-Reranker |
## 更新日志 ## 更新日志
### 2026-05-29
- generate代码重构完成
### 2026-05-28 ### 2026-05-28
- generate代码重构进度 1/3 - generate代码重构进度 1/3
@@ -53,9 +56,6 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
### 2026-05-09 ### 2026-05-09
- 合并 pr/eastgold15/46, 添加 aha-ui - 合并 pr/eastgold15/46, 添加 aha-ui
### 2026-04-25
- VoxCPM 更新流式生成
**[查看完整更新日志](docs/changelog.zh-CN.md)** → **[查看完整更新日志](docs/changelog.zh-CN.md)** →
+3
View File
@@ -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/), 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). and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
### 2026-05-29
- generate code refactored
### 2026-05-28 ### 2026-05-28
- generate code refactoring progress 1/3 - generate code refactoring progress 1/3
+3
View File
@@ -5,6 +5,9 @@
格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/) 格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/)
本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。 本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。
### 2026-05-29
- generate代码重构完成
### 2026-05-28 ### 2026-05-28
- generate代码重构进度 1/3 - generate代码重构进度 1/3
+22 -73
View File
@@ -3,18 +3,16 @@ use std::collections::HashMap;
use crate::{ use crate::{
models::common::{ models::common::{
MultiModalData, MultiModalData,
generate::{GenerationContext, generate_generic, generate_stream_generic}, generate::{GenerationDataProvider, PrepareData},
}, },
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, params::chat::ChatCompletionParameters,
}; };
use anyhow::{Result, anyhow}; use anyhow::{Result, anyhow};
use candle_core::{DType, Device, pickle::read_all_with_key}; use candle_core::{DType, Device, pickle::read_all_with_key};
use candle_nn::VarBuilder; use candle_nn::VarBuilder;
use rocket::futures::Stream;
use crate::{ use crate::{
models::{ models::{
GenerateModel,
fun_asr_nano::{ fun_asr_nano::{
config::FunASRNanoConfig, model::FunAsrNanoModel, processor::FunAsrNanoProcessor, config::FunASRNanoConfig, model::FunAsrNanoModel, processor::FunAsrNanoProcessor,
}, },
@@ -94,79 +92,30 @@ impl FunAsrNanoGenerateModel {
} }
} }
impl GenerateModel for FunAsrNanoGenerateModel { impl GenerationDataProvider for FunAsrNanoGenerateModel {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> { fn get_temperature(&self, req_temp: Option<f32>) -> Option<f32> {
let temperature = mes Some(req_temp.unwrap_or(self.generation_config.temperature))
.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,
)
} }
fn generate_stream( fn get_top_p(&self, req_top_p: Option<f32>) -> Option<f32> {
&mut self, Some(req_top_p.unwrap_or(self.generation_config.top_p))
mes: ChatCompletionParameters, }
) -> Result<
Box< fn get_top_k(&self, top_k: Option<usize>) -> Option<usize> {
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>> Some(top_k.unwrap_or(self.generation_config.top_k))
+ Send }
+ Unpin fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData> {
+ '_, let (speech, fbank_mask, input_ids) = self.processor.process_info(mes, &self.tokenizer)?;
>,
> {
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 speech = speech.to_dtype(self.dtype)?;
let data_vec = vec![speech.into(), fbank_mask.into()]; let data_vec = vec![speech.into(), fbank_mask.into()];
let data = MultiModalData::new(data_vec); let multi_model_data = MultiModalData::new(data_vec);
let stream = generate_stream_generic(
&mut self.model, Ok(PrepareData {
&self.tokenizer, in_reasoning: false,
input_ids, input_ids,
data, multi_model_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)))
} }
} }
crate::impl_generate_model!(FunAsrNanoGenerateModel);
+14 -73
View File
@@ -1,23 +1,19 @@
use crate::{ use crate::{
models::common::{ models::common::{
MultiModalData, MultiModalData,
generate::{GenerationContext, generate_generic, generate_stream_generic}, generate::{GenerationDataProvider, PrepareData},
}, },
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, params::chat::ChatCompletionParameters,
}; };
use anyhow::Result; use anyhow::Result;
use candle_core::{DType, Device, Tensor}; use candle_core::{DType, Device, Tensor};
use candle_nn::VarBuilder; use candle_nn::VarBuilder;
use rocket::futures::Stream;
use crate::{ use crate::{
chat_template::ChatTemplate, chat_template::ChatTemplate,
models::{ models::glm_asr_nano::{
GenerateModel,
glm_asr_nano::{
config::GlmAsrNanoConfig, model::GlmAsrNanoModel, processor::GlmAsrNanoProcessor, config::GlmAsrNanoConfig, model::GlmAsrNanoModel, processor::GlmAsrNanoProcessor,
}, },
},
tokenizer::TokenizerModel, tokenizer::TokenizerModel,
utils::{find_type_files, get_device, get_dtype}, utils::{find_type_files, get_device, get_dtype},
}; };
@@ -63,77 +59,22 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
} }
} }
impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> { impl<'a> GenerationDataProvider for GlmAsrNanoGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> { fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData> {
let seed = mes.seed.unwrap_or(34562) as u64; let render_text: String = self.chat_template.apply_chat_template(mes)?;
let render_text: String = self.chat_template.apply_chat_template(&mes)?;
let (input_features, audio_token_lengths, replace_text) = 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_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
let input_features = input_features.to_dtype(self.dtype)?; let input_features = input_features.to_dtype(self.dtype)?;
let audio_token_lengths = Tensor::new(audio_token_lengths, &self.device)?; 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_vec = vec![input_features.into(), audio_token_lengths.into()];
let data = MultiModalData::new(data_vec); let multi_model_data = MultiModalData::new(data_vec);
generate_generic( Ok(PrepareData {
&mut self.model, in_reasoning: false,
&self.tokenizer,
input_ids, input_ids,
data, multi_model_data,
&mut ctx, })
&self.model_name,
)
}
fn generate_stream(
&mut self,
mes: ChatCompletionParameters,
) -> Result<
Box<
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ 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)))
} }
} }
crate::impl_generate_model!(GlmAsrNanoGenerateModel<'a>);
+17 -105
View File
@@ -2,33 +2,26 @@
use crate::{ use crate::{
models::common::{ models::common::{
MultiModalData, 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::{ use crate::{
// chat_template::ChatTemplate, models::glm_ocr::{
models::{
GenerateModel,
glm_ocr::{
config::{GlmOcrConfig, GlmOcrGenerationConfig}, config::{GlmOcrConfig, GlmOcrGenerationConfig},
model::GlmOcrModel, model::GlmOcrModel,
processor::GlmOcrProcessor, processor::GlmOcrProcessor,
}, },
},
tokenizer::TokenizerModel, tokenizer::TokenizerModel,
utils::{ utils::{
extract_user_text, find_type_files, get_device, get_dtype, img_utils::extract_image_url, 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 { pub struct GlmOcrGenerateModel {
// chat_template: ChatTemplate<'a>,
tokenizer: TokenizerModel, tokenizer: TokenizerModel,
processor: GlmOcrProcessor, processor: GlmOcrProcessor,
model: GlmOcrModel, model: GlmOcrModel,
@@ -44,7 +37,6 @@ pub struct GlmOcrGenerateModel {
impl GlmOcrGenerateModel { impl GlmOcrGenerateModel {
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> { pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
// let chat_template = ChatTemplate::init(path)?;
let tokenizer = TokenizerModel::init(path)?; let tokenizer = TokenizerModel::init(path)?;
let config_path = path.to_string() + "/config.json"; let config_path = path.to_string() + "/config.json";
let cfg: GlmOcrConfig = serde_json::from_slice(&std::fs::read(config_path)?)?; let cfg: GlmOcrConfig = serde_json::from_slice(&std::fs::read(config_path)?)?;
@@ -64,7 +56,6 @@ impl GlmOcrGenerateModel {
.unwrap_or("ZhipuAI/GLM-OCR") .unwrap_or("ZhipuAI/GLM-OCR")
.to_string(); .to_string();
Ok(Self { Ok(Self {
// chat_template,
tokenizer, tokenizer,
processor, processor,
model, model,
@@ -80,17 +71,15 @@ impl GlmOcrGenerateModel {
} }
} }
impl GenerateModel for GlmOcrGenerateModel { impl GenerationDataProvider for GlmOcrGenerateModel {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> { fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData> {
let seed = mes.seed.unwrap_or(34562) as u64; let image_urls = extract_image_url(mes);
// Extract image path and prompt from messages
let image_urls = extract_image_url(&mes);
let image_path = image_urls let image_path = image_urls
.first() .first()
.ok_or_else(|| anyhow!("No image provided"))?; .ok_or_else(|| anyhow!("No image provided"))?;
// Get prompt text from messages // Get prompt text from messages
let mut prompt = extract_user_text(&mes)?; let mut prompt = extract_user_text(mes)?;
if prompt.chars().count() == 0 { if prompt.chars().count() == 0 {
prompt = "Extract all text from this image.".to_string() prompt = "Extract all text from this image.".to_string()
} }
@@ -108,95 +97,18 @@ impl GenerateModel for GlmOcrGenerateModel {
)?; )?;
let input_ids = processed.input_ids; 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![ let data_vec = vec![
processed.pixel_values.into(), processed.pixel_values.into(),
processed.grid_thw.into(), processed.grid_thw.into(),
processed.image_mask.into(), processed.image_mask.into(),
]; ];
let data = MultiModalData::new(data_vec); let multi_model_data = MultiModalData::new(data_vec);
generate_generic( Ok(PrepareData {
&mut self.model, in_reasoning: false,
&self.tokenizer,
input_ids, input_ids,
data, multi_model_data,
&mut ctx, })
&self.model_name,
)
}
fn generate_stream(
&mut self,
mes: ChatCompletionParameters,
) -> Result<
Box<
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ 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)))
} }
} }
crate::impl_generate_model!(GlmOcrGenerateModel);
+24 -86
View File
@@ -1,25 +1,21 @@
use crate::{ use crate::{
models::common::{ models::common::{
MultiModalData, MultiModalData,
generate::{GenerationContext, generate_generic, generate_stream_generic}, generate::{GenerationDataProvider, PrepareData},
}, },
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, params::chat::ChatCompletionParameters,
}; };
use anyhow::Result; use anyhow::Result;
use candle_core::{DType, Device}; use candle_core::{DType, Device};
use candle_nn::VarBuilder; use candle_nn::VarBuilder;
use rocket::futures::Stream;
use crate::{ use crate::{
chat_template::ChatTemplate, chat_template::ChatTemplate,
models::{ models::hunyuan_ocr::{
GenerateModel,
hunyuan_ocr::{
config::{HunYuanVLConfig, HunyuanOCRGenerationConfig}, config::{HunYuanVLConfig, HunyuanOCRGenerationConfig},
model::HunyuanVLModel, model::HunyuanVLModel,
processor::HunyuanVLProcessor, processor::HunyuanVLProcessor,
}, },
},
tokenizer::TokenizerModel, tokenizer::TokenizerModel,
utils::{find_type_files, get_device, get_dtype}, utils::{find_type_files, get_device, get_dtype},
}; };
@@ -68,72 +64,24 @@ impl<'a> HunyuanOCRGenerateModel<'a> {
} }
} }
impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> { impl<'a> GenerationDataProvider for HunyuanOCRGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> { fn get_temperature(&self, req_temp: Option<f32>) -> Option<f32> {
let temperature = mes Some(req_temp.unwrap_or(self.generation_config.temperature))
.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,
)
} }
fn generate_stream( fn get_top_p(&self, req_top_p: Option<f32>) -> Option<f32> {
&mut self, Some(req_top_p.unwrap_or(self.generation_config.top_p))
mes: ChatCompletionParameters, }
) -> Result<
Box< fn get_top_k(&self, top_k: Option<usize>) -> Option<usize> {
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>> Some(top_k.unwrap_or(self.generation_config.top_k))
+ Send }
+ Unpin
+ '_, fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData> {
>, let mes_render = self.chat_template.apply_chat_template(mes)?;
> {
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 let data = self
.pre_processor .pre_processor
.process_info(&mes, &self.tokenizer, &mes_render)?; .process_info(mes, &self.tokenizer, &mes_render)?;
let sample_len = mes.max_tokens.unwrap_or(1024);
let input_ids = data.input_ids; let input_ids = data.input_ids;
let data_vec = vec![ let data_vec = vec![
data.pixel_values, data.pixel_values,
@@ -141,23 +89,13 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
data.image_mask.into(), data.image_mask.into(),
data.position_ids.into(), data.position_ids.into(),
]; ];
let data = MultiModalData::new(data_vec); let multi_model_data = MultiModalData::new(data_vec);
let stream = generate_stream_generic( Ok(PrepareData {
&mut self.model, in_reasoning: false,
&self.tokenizer,
input_ids, input_ids,
data, multi_model_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)))
} }
} }
crate::impl_generate_model!(HunyuanOCRGenerateModel<'a>);
+13 -72
View File
@@ -1,9 +1,9 @@
use crate::{ use crate::{
models::common::{ models::common::{
MultiModalData, MultiModalData,
generate::{GenerationContext, generate_generic, generate_stream_generic}, generate::{GenerationDataProvider, PrepareData},
}, },
params::chat::{ChatCompletionParameters, ChatCompletionResponse}, params::chat::ChatCompletionParameters,
}; };
use anyhow::Result; use anyhow::Result;
use candle_core::{DType, Device}; use candle_core::{DType, Device};
@@ -12,7 +12,6 @@ use candle_nn::VarBuilder;
use crate::{ use crate::{
chat_template::ChatTemplate, chat_template::ChatTemplate,
models::{ models::{
GenerateModel,
lfm2::config::Lfm2GenerateConfig, lfm2::config::Lfm2GenerateConfig,
lfm2vl::{config::Lfm2VLConfig, model::Lfm2VLModel, processor::Lfm2VLProcessor}, lfm2vl::{config::Lfm2VLConfig, model::Lfm2VLModel, processor::Lfm2VLProcessor},
}, },
@@ -60,82 +59,24 @@ impl<'a> Lfm2VLGenerateModel<'a> {
} }
} }
impl<'a> GenerateModel for Lfm2VLGenerateModel<'a> { impl<'a> GenerationDataProvider for Lfm2VLGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> { fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData> {
let mes_render = self.chat_template.apply_chat_template(&mes)?; let mes_render = self.chat_template.apply_chat_template(mes)?;
let (pixel_values, pixel_attention_mask, spatial_shapes, text) = 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 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![ let data_vec = vec![
pixel_values.into(), pixel_values.into(),
pixel_attention_mask.into(), pixel_attention_mask.into(),
spatial_shapes.into(), spatial_shapes.into(),
]; ];
let data = MultiModalData::new(data_vec); let multi_model_data = MultiModalData::new(data_vec);
generate_generic( Ok(PrepareData {
&mut self.model, in_reasoning: false,
&self.tokenizer,
input_ids, input_ids,
data, multi_model_data,
&mut ctx, })
&self.model_name,
)
}
fn generate_stream(
&mut self,
mes: ChatCompletionParameters,
) -> Result<
Box<
dyn rocket::futures::Stream<
Item = Result<crate::params::chat::ChatCompletionChunkResponse, anyhow::Error>,
> + 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)))
} }
} }
crate::impl_generate_model!(Lfm2VLGenerateModel<'a>);
+14 -84
View File
@@ -1,21 +1,16 @@
use crate::models::common::MultiModalData; use crate::models::common::MultiModalData;
use crate::models::common::generate::{ use crate::models::common::generate::{GenerationDataProvider, PrepareData};
GenerationContext, generate_generic, generate_stream_generic, use crate::params::chat::ChatCompletionParameters;
};
use crate::params::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
use anyhow::Result; use anyhow::Result;
use candle_core::{D, DType, Device, IndexOp, Tensor}; use candle_core::{D, DType, Device, IndexOp, Tensor};
use candle_nn::VarBuilder; use candle_nn::VarBuilder;
use rocket::futures::Stream;
use crate::models::paddleocr_vl::config::{PaddleOCRVLConfig, PaddleOCRVLPreprocessorConfig}; use crate::models::paddleocr_vl::config::{PaddleOCRVLConfig, PaddleOCRVLPreprocessorConfig};
use crate::models::paddleocr_vl::model::PaddleOCRVLModel; use crate::models::paddleocr_vl::model::PaddleOCRVLModel;
use crate::models::paddleocr_vl::processor::PaddleOCRVLProcessor; use crate::models::paddleocr_vl::processor::PaddleOCRVLProcessor;
use crate::utils::tensor_utils::get_equal_mask; use crate::utils::tensor_utils::get_equal_mask;
use crate::utils::{find_type_files, get_device, get_dtype}; 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> { pub struct PaddleOCRVLGenerateModel<'a> {
chat_template: ChatTemplate<'a>, chat_template: ChatTemplate<'a>,
@@ -60,11 +55,11 @@ impl<'a> PaddleOCRVLGenerateModel<'a> {
} }
} }
impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> { impl<'a> GenerationDataProvider for PaddleOCRVLGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> { fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData> {
let mes_render = self.chat_template.apply_chat_template(&mes)?; let mes_render = self.chat_template.apply_chat_template(mes)?;
let (replace_text, pixel_values, image_grid_thw) = 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 input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
let image_mask = get_equal_mask(&input_ids, self.cfg.image_token_id)?; 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)? .cumsum(D::Minus1)?
.to_dtype(candle_core::DType::U32)? .to_dtype(candle_core::DType::U32)?
.broadcast_sub(&Tensor::new(vec![1_u32], input_ids.device())?)?; .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![ let data_vec = vec![
pixel_values, pixel_values,
image_grid_thw, image_grid_thw,
image_mask.into(), image_mask.into(),
cache_position.into(), cache_position.into(),
]; ];
let data = MultiModalData::new(data_vec); let multi_model_data = MultiModalData::new(data_vec);
generate_generic( Ok(PrepareData {
&mut self.model, in_reasoning: false,
&self.tokenizer,
input_ids, input_ids,
data, multi_model_data,
&mut ctx, })
&self.model_name,
)
}
fn generate_stream(
&mut self,
mes: ChatCompletionParameters,
) -> Result<
Box<
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ 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)))
} }
} }
crate::impl_generate_model!(PaddleOCRVLGenerateModel<'a>);
+24 -83
View File
@@ -2,11 +2,11 @@ use crate::{
models::common::{ models::common::{
MultiModalData, MultiModalData,
generate::{ generate::{
GenerationContext, generate_generic, generate_generic_text, generate_stream_generic, GenerationContext, GenerationDataProvider, PrepareData, generate_generic_text,
generate_stream_generic_text, generate_stream_generic_text,
}, },
}, },
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, params::chat::ChatCompletionParameters,
}; };
use anyhow::{Result, anyhow}; use anyhow::{Result, anyhow};
use candle_core::{DType, Device, quantized::gguf_file}; use candle_core::{DType, Device, quantized::gguf_file};
@@ -16,7 +16,6 @@ use rocket::futures::Stream;
use crate::{ use crate::{
chat_template::ChatTemplate, chat_template::ChatTemplate,
models::{ models::{
GenerateModel,
common::gguf::Gguf, common::gguf::Gguf,
qwen3_5::{config::Qwen3_5Config, model::Qwen3_5Model}, qwen3_5::{config::Qwen3_5Config, model::Qwen3_5Model},
qwen3vl::processor::Qwen3VLProcessor, qwen3vl::processor::Qwen3VLProcessor,
@@ -247,71 +246,25 @@ impl<'a> Qwen3_5GenerateModel<'a> {
} }
} }
impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> { impl<'a> GenerationDataProvider for Qwen3_5GenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> { fn get_temperature(&self, req_temp: Option<f32>) -> Option<f32> {
let seed = mes.seed.unwrap_or(32768) as u64; Some(req_temp.unwrap_or(0.4))
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,
)
} }
fn generate_stream( fn get_top_p(&self, req_top_p: Option<f32>) -> Option<f32> {
&mut self, Some(req_top_p.unwrap_or(0.95))
mes: ChatCompletionParameters, }
) -> Result<
Box< fn get_top_k(&self, top_k: Option<usize>) -> Option<usize> {
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>> Some(top_k.unwrap_or(40))
+ Send }
+ Unpin
+ '_, fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData> {
>, let mes_render = self.chat_template.apply_chat_template(mes)?;
> { let in_reasoning = self.is_in_reasoning(&mes_render);
let mes_render = self.chat_template.apply_chat_template(&mes)?;
let in_reasoning = mes_render.ends_with("<think>\n");
let (mes_text, pixel_values, image_grid_thw, pixel_values_video, video_grid_thw) = let (mes_text, pixel_values, image_grid_thw, pixel_values_video, video_grid_thw) =
if let Some(processor) = &self.pre_processor { 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.replace_text,
input.pixel_values, input.pixel_values,
@@ -323,31 +276,19 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
(mes_render, None, None, None, None) (mes_render, None, None, None, None)
}; };
let input_ids = self.tokenizer.text_encode(mes_text, &self.device)?; let input_ids = self.tokenizer.text_encode(mes_text, &self.device)?;
let sample_len = mes.max_tokens.unwrap_or(1024);
let data_vec = vec![ let data_vec = vec![
pixel_values, pixel_values,
image_grid_thw, image_grid_thw,
pixel_values_video, pixel_values_video,
video_grid_thw, video_grid_thw,
]; ];
let data = MultiModalData::new(data_vec); let multi_model_data = MultiModalData::new(data_vec);
let seed = mes.seed.unwrap_or(34562) as u64; Ok(PrepareData {
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,
in_reasoning, in_reasoning,
&self.device, input_ids,
&self.model_name, multi_model_data,
)?; })
Ok(Box::new(Box::pin(stream)))
} }
} }
crate::impl_generate_model!(Qwen3_5GenerateModel<'a>);
+24 -85
View File
@@ -1,19 +1,17 @@
use crate::{ use crate::{
models::common::{ models::common::{
MultiModalData, MultiModalData,
generate::{GenerationContext, generate_generic, generate_stream_generic}, generate::{GenerationDataProvider, PrepareData},
}, },
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, params::chat::ChatCompletionParameters,
}; };
use anyhow::Result; use anyhow::Result;
use candle_core::{DType, Device, Tensor}; use candle_core::{DType, Device, Tensor};
use candle_nn::VarBuilder; use candle_nn::VarBuilder;
use rocket::futures::Stream;
use crate::{ use crate::{
chat_template::ChatTemplate, chat_template::ChatTemplate,
models::{ models::{
GenerateModel,
qwen3::config::Qwen3GenerationConfig, qwen3::config::Qwen3GenerationConfig,
qwen3vl::{config::Qwen3VLConfig, model::Qwen3VLModel, processor::Qwen3VLProcessor}, qwen3vl::{config::Qwen3VLConfig, model::Qwen3VLModel, processor::Qwen3VLProcessor},
}, },
@@ -65,76 +63,28 @@ impl<'a> Qwen3VLGenerateModel<'a> {
} }
} }
impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { impl<'a> GenerationDataProvider for Qwen3VLGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> { fn get_temperature(&self, req_temp: Option<f32>) -> Option<f32> {
let temperature = mes Some(req_temp.unwrap_or(self.generation_config.temperature))
.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,
)
} }
fn generate_stream( fn get_top_p(&self, req_top_p: Option<f32>) -> Option<f32> {
&mut self, Some(req_top_p.unwrap_or(self.generation_config.top_p))
mes: ChatCompletionParameters, }
) -> Result<
Box< fn get_top_k(&self, top_k: Option<usize>) -> Option<usize> {
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>> Some(top_k.unwrap_or(self.generation_config.top_k))
+ Send }
+ Unpin
+ '_, fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData> {
>, let mes_render = self.chat_template.apply_chat_template(mes)?;
> { let in_reasoning = self.is_in_reasoning(&mes_render);
let temperature = mes let input = self.pre_processor.process_info(mes, &mes_render)?;
.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("<think>\n");
let input = self.pre_processor.process_info(&mes, &mes_render)?;
let input_ids = self let input_ids = self
.tokenizer .tokenizer
.text_encode(input.replace_text.clone(), &self.device)?; .text_encode(input.replace_text.clone(), &self.device)?;
let seq_len = input_ids.dim(1)?; let seq_len = input_ids.dim(1)?;
let cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?; 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![ let data_vec = vec![
input.pixel_values, input.pixel_values,
input.image_grid_thw, input.image_grid_thw,
@@ -142,24 +92,13 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
input.video_grid_thw, input.video_grid_thw,
cache_position.into(), cache_position.into(),
]; ];
let data = MultiModalData::new(data_vec); let multi_model_data = MultiModalData::new(data_vec);
let seed = mes.seed.unwrap_or(34562) as u64; Ok(PrepareData {
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,
in_reasoning, in_reasoning,
&self.device, input_ids,
&self.model_name, multi_model_data,
)?; })
Ok(Box::new(Box::pin(stream)))
} }
} }
crate::impl_generate_model!(Qwen3VLGenerateModel<'a>);
+1 -1
View File
@@ -21,7 +21,7 @@ fn fun_asr_nano_generate() -> Result<()> {
"type": "audio", "type": "audio",
"audio_url": "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"
} }
}, },
{ {
+1 -1
View File
@@ -22,7 +22,7 @@ fn glm_asr_nano_generate() -> Result<()> {
"type": "audio", "type": "audio",
"audio_url": "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"
} }
}, },
{ {
+1 -1
View File
@@ -52,7 +52,7 @@ fn glm_ocr_generate() -> Result<()> {
#[tokio::test] #[tokio::test]
async fn glm_ocr_stream() -> Result<()> { 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#" let message = r#"
{ {
+1 -1
View File
@@ -49,7 +49,7 @@ fn hunyuan_ocr_generate() -> Result<()> {
#[tokio::test] #[tokio::test]
async fn hunyuan_ocr_stream() -> Result<()> { 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#" let message = r#"
{ {
+1 -1
View File
@@ -60,7 +60,7 @@ async fn lfm2vl_stream() -> Result<()> {
let save_dir = let save_dir =
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get 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-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#" let message = r#"
{ {
"model": "lfm2vl", "model": "lfm2vl",
+5 -5
View File
@@ -104,7 +104,7 @@ fn qwen3vl_generate() -> Result<()> {
#[tokio::test] #[tokio::test]
async fn qwen3vl_stream() -> Result<()> { 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 = let save_dir =
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get 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", "role": "user",
"content": [ "content": [
{ {
"type": "video", "type": "image",
"video_url": "image_url":
{ {
"url": "./assets/video/video_test.mp4" "url": "file://./assets/img/ocr_test1.png"
} }
}, },
{ {
"type": "text", "type": "text",
"text": "视频中发生了什么?" "text": "OCR"
} }
] ]
} }