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 |
|
| **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
@@ -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)** →
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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);
|
||||||
|
|||||||
@@ -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
@@ -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);
|
||||||
|
|||||||
@@ -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>);
|
||||||
|
|||||||
@@ -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>);
|
||||||
|
|||||||
@@ -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>);
|
||||||
|
|||||||
@@ -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>);
|
||||||
|
|||||||
@@ -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>);
|
||||||
|
|||||||
@@ -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"
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -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"
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -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#"
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -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#"
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
@@ -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"
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user