generate code refactored
This commit is contained in:
@@ -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?
|
||||
|
||||
+3
-3
@@ -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)** →
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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<ChatCompletionResponse> {
|
||||
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<f32>) -> Option<f32> {
|
||||
Some(req_temp.unwrap_or(self.generation_config.temperature))
|
||||
}
|
||||
|
||||
fn generate_stream(
|
||||
&mut self,
|
||||
mes: ChatCompletionParameters,
|
||||
) -> Result<
|
||||
Box<
|
||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
||||
+ 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<f32>) -> Option<f32> {
|
||||
Some(req_top_p.unwrap_or(self.generation_config.top_p))
|
||||
}
|
||||
|
||||
fn get_top_k(&self, top_k: Option<usize>) -> Option<usize> {
|
||||
Some(top_k.unwrap_or(self.generation_config.top_k))
|
||||
}
|
||||
fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData> {
|
||||
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);
|
||||
|
||||
@@ -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<ChatCompletionResponse> {
|
||||
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<PrepareData> {
|
||||
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<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)))
|
||||
multi_model_data,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
crate::impl_generate_model!(GlmAsrNanoGenerateModel<'a>);
|
||||
|
||||
+20
-108
@@ -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<DType>) -> Result<Self> {
|
||||
// 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<ChatCompletionResponse> {
|
||||
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<PrepareData> {
|
||||
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<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)))
|
||||
multi_model_data,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
crate::impl_generate_model!(GlmOcrGenerateModel);
|
||||
|
||||
@@ -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<ChatCompletionResponse> {
|
||||
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<f32>) -> Option<f32> {
|
||||
Some(req_temp.unwrap_or(self.generation_config.temperature))
|
||||
}
|
||||
|
||||
fn generate_stream(
|
||||
&mut self,
|
||||
mes: ChatCompletionParameters,
|
||||
) -> Result<
|
||||
Box<
|
||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
||||
+ 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<f32>) -> Option<f32> {
|
||||
Some(req_top_p.unwrap_or(self.generation_config.top_p))
|
||||
}
|
||||
|
||||
fn get_top_k(&self, top_k: Option<usize>) -> Option<usize> {
|
||||
Some(top_k.unwrap_or(self.generation_config.top_k))
|
||||
}
|
||||
|
||||
fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData> {
|
||||
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>);
|
||||
|
||||
@@ -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<ChatCompletionResponse> {
|
||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||
impl<'a> GenerationDataProvider for Lfm2VLGenerateModel<'a> {
|
||||
fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData> {
|
||||
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<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)))
|
||||
multi_model_data,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
crate::impl_generate_model!(Lfm2VLGenerateModel<'a>);
|
||||
|
||||
@@ -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<ChatCompletionResponse> {
|
||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||
impl<'a> GenerationDataProvider for PaddleOCRVLGenerateModel<'a> {
|
||||
fn get_data(&self, mes: &ChatCompletionParameters) -> Result<PrepareData> {
|
||||
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<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)))
|
||||
multi_model_data,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
crate::impl_generate_model!(PaddleOCRVLGenerateModel<'a>);
|
||||
|
||||
@@ -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<ChatCompletionResponse> {
|
||||
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<f32>) -> Option<f32> {
|
||||
Some(req_temp.unwrap_or(0.4))
|
||||
}
|
||||
|
||||
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 in_reasoning = mes_render.ends_with("<think>\n");
|
||||
fn get_top_p(&self, req_top_p: Option<f32>) -> Option<f32> {
|
||||
Some(req_top_p.unwrap_or(0.95))
|
||||
}
|
||||
|
||||
fn get_top_k(&self, top_k: Option<usize>) -> Option<usize> {
|
||||
Some(top_k.unwrap_or(40))
|
||||
}
|
||||
|
||||
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_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>);
|
||||
|
||||
@@ -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<ChatCompletionResponse> {
|
||||
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<f32>) -> Option<f32> {
|
||||
Some(req_temp.unwrap_or(self.generation_config.temperature))
|
||||
}
|
||||
|
||||
fn generate_stream(
|
||||
&mut self,
|
||||
mes: ChatCompletionParameters,
|
||||
) -> Result<
|
||||
Box<
|
||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
||||
+ 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("<think>\n");
|
||||
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||
fn get_top_p(&self, req_top_p: Option<f32>) -> Option<f32> {
|
||||
Some(req_top_p.unwrap_or(self.generation_config.top_p))
|
||||
}
|
||||
|
||||
fn get_top_k(&self, top_k: Option<usize>) -> Option<usize> {
|
||||
Some(top_k.unwrap_or(self.generation_config.top_k))
|
||||
}
|
||||
|
||||
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 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>);
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
},
|
||||
{
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
},
|
||||
{
|
||||
|
||||
@@ -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#"
|
||||
{
|
||||
|
||||
@@ -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#"
|
||||
{
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user