refactor generate code
This commit is contained in:
@@ -29,11 +29,11 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
|
|||||||
|
|
||||||
| Category | Models |
|
| Category | Models |
|
||||||
|----------|--------|
|
|----------|--------|
|
||||||
| **Text** | Qwen3, MiniCPM4, <br> LFM2, LFM2.5 |
|
| **Text** | Qwen3, MiniCPM4, LFM2, LFM2.5 |
|
||||||
| **Vision** | Qwen2.5-VL, Qwen3-VL, Qwen3.5, <br> LFM2.5-VL, LFM2-VL |
|
| **Vision** | Qwen2.5-VL, Qwen3-VL, Qwen3.5, <br> LFM2.5-VL, LFM2-VL |
|
||||||
| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2 , <br> PaddleOCR-VL, PaddleOCR-VL1.5, <br> Hunyuan-OCR, GLM-OCR |
|
| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2 , PaddleOCR-VL <br> PaddleOCR-VL1.5, Hunyuan-OCR, GLM-OCR |
|
||||||
| **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR |
|
| **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR |
|
||||||
| **Audio** | VoxCPM, VoxCPM1.5 |
|
| **TTS** | VoxCPM, VoxCPM1.5 |
|
||||||
| **Image** | RMBG-2.0 (background removal) |
|
| **Image** | RMBG-2.0 (background removal) |
|
||||||
|
|
||||||
## Why aha?
|
## Why aha?
|
||||||
@@ -46,6 +46,11 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
|
|||||||
- **🧠 Attention Optimization** - Optional Flash Attention support for optimized long sequence processing
|
- **🧠 Attention Optimization** - Optional Flash Attention support for optimized long sequence processing
|
||||||
|
|
||||||
## Changelog
|
## Changelog
|
||||||
|
### 2026-04-02
|
||||||
|
- refactor generate code
|
||||||
|
- \<think\>...\</think\> The content of the thought chain is returned using the reasoning_content field.
|
||||||
|
- chat response add time info
|
||||||
|
|
||||||
### 2026-04-01
|
### 2026-04-01
|
||||||
- refactor deepseek_ocr/fun_asr_nano generate code
|
- refactor deepseek_ocr/fun_asr_nano generate code
|
||||||
|
|
||||||
@@ -64,9 +69,6 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
|
|||||||
- add LFM2.5-1.2B-Instruct
|
- add LFM2.5-1.2B-Instruct
|
||||||
- add LFM2-1.2B
|
- add LFM2-1.2B
|
||||||
|
|
||||||
### v0.2.3 (2026-03-18)
|
|
||||||
- add DeepSeek-OCR-2
|
|
||||||
|
|
||||||
**[View full changelog](docs/changelog.md)** →
|
**[View full changelog](docs/changelog.md)** →
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+9
-3
@@ -28,11 +28,11 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
|
|||||||
|
|
||||||
| 类别 | 模型 |
|
| 类别 | 模型 |
|
||||||
|------|------|
|
|------|------|
|
||||||
| **文本** | Qwen3, MiniCPM4, <br> LFM2, LFM2.5 |
|
| **文本** | Qwen3, MiniCPM4, LFM2, LFM2.5 |
|
||||||
| **视觉** | Qwen2.5-VL, Qwen3-VL, Qwen3.5 <br> LFM2.5-VL, LFM2-VL |
|
| **视觉** | Qwen2.5-VL, Qwen3-VL, Qwen3.5 <br> LFM2.5-VL, LFM2-VL |
|
||||||
| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2 , <br> PaddleOCR-VL, PaddleOCR-VL1.5, <br> Hunyuan-OCR, GLM-OCR |
|
| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2, PaddleOCR-VL, <br>PaddleOCR-VL1.5, Hunyuan-OCR, GLM-OCR |
|
||||||
| **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR |
|
| **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR |
|
||||||
| **音频** | VoxCPM, VoxCPM1.5 |
|
| **TTS** | VoxCPM, VoxCPM1.5 |
|
||||||
| **图像** | RMBG-2.0 (背景移除) |
|
| **图像** | RMBG-2.0 (背景移除) |
|
||||||
|
|
||||||
## 为什么选择 aha?
|
## 为什么选择 aha?
|
||||||
@@ -45,6 +45,12 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
|
|||||||
- **🧠 注意力优化** - 可选 Flash Attention 支持,优化长序列处理
|
- **🧠 注意力优化** - 可选 Flash Attention 支持,优化长序列处理
|
||||||
|
|
||||||
## 更新日志
|
## 更新日志
|
||||||
|
## Changelog
|
||||||
|
### 2026-04-02
|
||||||
|
- 重构生成代码
|
||||||
|
- \<think\>...\</think\> 思维链内容使用reasoning_content字段返回。
|
||||||
|
- 对话返回添加耗时信息
|
||||||
|
|
||||||
### 2026-04-01
|
### 2026-04-01
|
||||||
- 重构 deepseek_ocr/fun_asr_nano 生成代码
|
- 重构 deepseek_ocr/fun_asr_nano 生成代码
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,11 @@ 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-04-02
|
||||||
|
- refactor generate code
|
||||||
|
- \<think\>...\</think\> The content of the thought chain is returned using the reasoning_content field.
|
||||||
|
- response add time info
|
||||||
|
|
||||||
### 2026-04-01
|
### 2026-04-01
|
||||||
- refactor deepseek_ocr/fun_asr_nano generate code
|
- refactor deepseek_ocr/fun_asr_nano generate code
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,11 @@
|
|||||||
格式基于 [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-04-02
|
||||||
|
- 重构生成代码
|
||||||
|
- \<think\>...\</think\> 思维链内容使用reasoning_content字段返回。
|
||||||
|
- 对话返回添加耗时信息
|
||||||
|
|
||||||
### 2026-04-01
|
### 2026-04-01
|
||||||
- 重构 deepseek_ocr/fun_asr_nano 生成代码
|
- 重构 deepseek_ocr/fun_asr_nano 生成代码
|
||||||
|
|
||||||
|
|||||||
@@ -133,6 +133,9 @@ impl<'a> ChatTemplate<'a> {
|
|||||||
|
|
||||||
pub fn apply_chat_template(&self, messages: &ChatCompletionParameters) -> Result<String> {
|
pub fn apply_chat_template(&self, messages: &ChatCompletionParameters) -> Result<String> {
|
||||||
let enable_thinking = extract_metadata_value::<bool>(&messages.metadata, "enable_thinking");
|
let enable_thinking = extract_metadata_value::<bool>(&messages.metadata, "enable_thinking");
|
||||||
|
let mes_thinking_param = messages.enable_thinking;
|
||||||
|
let enable_thinking =
|
||||||
|
Some(enable_thinking.unwrap_or(false) || mes_thinking_param.unwrap_or(false));
|
||||||
let context = context! {
|
let context = context! {
|
||||||
messages => &messages.messages,
|
messages => &messages.messages,
|
||||||
tools => &messages.tools.as_ref(),
|
tools => &messages.tools.as_ref(),
|
||||||
|
|||||||
@@ -9,7 +9,10 @@ use crate::{
|
|||||||
models::common::{InferenceModel, MultiModalData},
|
models::common::{InferenceModel, MultiModalData},
|
||||||
params::chat::{ChatCompletionChunkResponse, ChatCompletionResponse},
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionResponse},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{build_completion_chunk_response, build_completion_response_with_time},
|
utils::response_utils::{
|
||||||
|
build_chunk_response_with_reasoning, build_chunk_response_with_usage,
|
||||||
|
build_completion_chunk_response, build_completion_response_with_time,
|
||||||
|
},
|
||||||
};
|
};
|
||||||
pub fn get_logit_processor(
|
pub fn get_logit_processor(
|
||||||
temperature: Option<f32>,
|
temperature: Option<f32>,
|
||||||
@@ -98,6 +101,17 @@ fn sample_and_push(
|
|||||||
Ok(token)
|
Ok(token)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TODO
|
||||||
|
// let logits = if self.repeat_penalty == 1. {
|
||||||
|
// logits
|
||||||
|
// } else {
|
||||||
|
// let start_at = generate.len().saturating_sub(self.repeat_last_n);
|
||||||
|
// candle_transformers::utils::apply_repeat_penalty(
|
||||||
|
// &logits,
|
||||||
|
// self.repeat_penalty,
|
||||||
|
// &generate[start_at..],
|
||||||
|
// )?
|
||||||
|
// };
|
||||||
pub fn generate_generic<M: InferenceModel>(
|
pub fn generate_generic<M: InferenceModel>(
|
||||||
model: &mut M,
|
model: &mut M,
|
||||||
tokenizer: &TokenizerModel,
|
tokenizer: &TokenizerModel,
|
||||||
@@ -155,6 +169,7 @@ pub fn generate_stream_generic<M: InferenceModel>(
|
|||||||
top_k: Option<usize>,
|
top_k: Option<usize>,
|
||||||
seed: u64,
|
seed: u64,
|
||||||
max_tokens: u32,
|
max_tokens: u32,
|
||||||
|
in_reasoning: bool,
|
||||||
device: &Device,
|
device: &Device,
|
||||||
model_name: &str,
|
model_name: &str,
|
||||||
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
||||||
@@ -167,14 +182,23 @@ pub fn generate_stream_generic<M: InferenceModel>(
|
|||||||
max_tokens,
|
max_tokens,
|
||||||
device.clone(),
|
device.clone(),
|
||||||
);
|
);
|
||||||
|
let prompt_tokens = ctx.seq_len as u32;
|
||||||
|
let mut prompt_secs = 0.0f64;
|
||||||
|
let mut completion_tokens = 0u32;
|
||||||
|
let mut completion_secs = 0.0f64;
|
||||||
let mut error_tokens = Vec::new();
|
let mut error_tokens = Vec::new();
|
||||||
let eos_ids = model.stop_token_ids();
|
let eos_ids = model.stop_token_ids();
|
||||||
let stream = stream! {
|
let stream = stream! {
|
||||||
let mut input_ids = input_ids;
|
let mut input_ids = input_ids;
|
||||||
|
let mut tool_call_id = None;
|
||||||
|
let mut tool_call_content = String::new();
|
||||||
|
let mut in_reasoning = in_reasoning;
|
||||||
// 处理 unicode 错误累积
|
// 处理 unicode 错误累积
|
||||||
for _ in 0..ctx.sample_len {
|
for _ in 0..ctx.sample_len {
|
||||||
|
let i_start = Instant::now();
|
||||||
let logits = if ctx.seqlen_offset == 0 {
|
let logits = if ctx.seqlen_offset == 0 {
|
||||||
model.forward_initial(&input_ids, ctx.seqlen_offset, data.clone())
|
model.forward_initial(&input_ids, ctx.seqlen_offset, data.clone())
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
model.forward_step(&input_ids, ctx.seqlen_offset)
|
model.forward_step(&input_ids, ctx.seqlen_offset)
|
||||||
}?;
|
}?;
|
||||||
@@ -183,6 +207,13 @@ pub fn generate_stream_generic<M: InferenceModel>(
|
|||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
ctx.logit_processor.sample(&logits)?
|
ctx.logit_processor.sample(&logits)?
|
||||||
};
|
};
|
||||||
|
completion_tokens += 1;
|
||||||
|
let i_duration = i_start.elapsed();
|
||||||
|
if ctx.seqlen_offset == 0 {
|
||||||
|
prompt_secs += i_duration.as_secs_f64();
|
||||||
|
} else {
|
||||||
|
completion_secs += i_duration.as_secs_f64();
|
||||||
|
};
|
||||||
|
|
||||||
// 解码(处理�的累积)
|
// 解码(处理�的累积)
|
||||||
let decode_ids = if error_tokens.is_empty() {
|
let decode_ids = if error_tokens.is_empty() {
|
||||||
@@ -204,9 +235,60 @@ pub fn generate_stream_generic<M: InferenceModel>(
|
|||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
error_tokens.clear();
|
error_tokens.clear();
|
||||||
yield Ok(build_completion_chunk_response(decoded, model_name, None, None));
|
if decoded.eq("<think>") {
|
||||||
|
in_reasoning = true;
|
||||||
|
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
if decoded.eq("</think>") {
|
||||||
|
in_reasoning = false;
|
||||||
|
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
// 处理特殊标记和工具调用
|
||||||
|
match decoded.as_str() {
|
||||||
|
"<tool_call>" => {
|
||||||
|
// 开始工具调用
|
||||||
|
tool_call_id = Some(uuid::Uuid::new_v4().to_string());
|
||||||
|
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
"</tool_call>" => {
|
||||||
|
// 结束工具调用
|
||||||
|
let chunk = build_completion_chunk_response(
|
||||||
|
decoded,
|
||||||
|
model_name,
|
||||||
|
tool_call_id.clone(),
|
||||||
|
Some(tool_call_content.clone())
|
||||||
|
);
|
||||||
|
tool_call_id = None;
|
||||||
|
tool_call_content = String::new();
|
||||||
|
yield Ok(chunk);
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
if tool_call_id.is_some() {
|
||||||
|
// 在工具调用过程中,收集工具调用内容
|
||||||
|
tool_call_content.push_str(&decoded);
|
||||||
|
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||||
|
continue;
|
||||||
|
} else {
|
||||||
|
|
||||||
|
// 正常文本输出
|
||||||
|
let chunk = if in_reasoning {
|
||||||
|
build_chunk_response_with_reasoning(decoded, model_name)
|
||||||
|
} else {
|
||||||
|
build_completion_chunk_response(
|
||||||
|
decoded, model_name,
|
||||||
|
None,
|
||||||
|
None
|
||||||
|
)};
|
||||||
|
yield Ok(chunk);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
if eos_ids.contains(&next_token) {
|
if eos_ids.contains(&next_token) {
|
||||||
|
yield Ok(build_chunk_response_with_usage(model_name, completion_tokens.into(), completion_secs.into(), prompt_tokens.into(), prompt_secs.into()));
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
input_ids = ctx.prepare_for_next_token(next_token)?;
|
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||||
|
|||||||
@@ -18,14 +18,18 @@ impl MultiModalData {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[allow(unused)]
|
||||||
pub trait InferenceModel {
|
pub trait InferenceModel {
|
||||||
/// 初始前向传播(考虑多模态输入)
|
/// 初始前向传播(考虑多模态输入)
|
||||||
|
/// 默认实现无特殊数据
|
||||||
fn forward_initial(
|
fn forward_initial(
|
||||||
&mut self,
|
&mut self,
|
||||||
input_ids: &Tensor,
|
input_ids: &Tensor,
|
||||||
seqlen_offset: usize,
|
seqlen_offset: usize,
|
||||||
data: MultiModalData,
|
data: MultiModalData,
|
||||||
) -> Result<Tensor>;
|
) -> Result<Tensor> {
|
||||||
|
Self::forward_step(self, input_ids, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
/// 后续前向传播(自回归步骤)
|
/// 后续前向传播(自回归步骤)
|
||||||
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor>;
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor>;
|
||||||
|
|||||||
@@ -171,6 +171,7 @@ impl GenerateModel for DeepseekOCRGenerateModel {
|
|||||||
None,
|
None,
|
||||||
seed,
|
seed,
|
||||||
max_tokens,
|
max_tokens,
|
||||||
|
false,
|
||||||
&self.device,
|
&self.device,
|
||||||
&self.model_name,
|
&self.model_name,
|
||||||
)?;
|
)?;
|
||||||
|
|||||||
@@ -30,8 +30,6 @@ pub struct FunAsrNanoGenerateModel {
|
|||||||
fun_asr_nano: FunAsrNanoModel,
|
fun_asr_nano: FunAsrNanoModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
dtype: DType,
|
dtype: DType,
|
||||||
// eos_token_id1: u32,
|
|
||||||
// eos_token_id2: u32,
|
|
||||||
generation_config: Qwen3GenerationConfig,
|
generation_config: Qwen3GenerationConfig,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
}
|
}
|
||||||
@@ -90,8 +88,6 @@ impl FunAsrNanoGenerateModel {
|
|||||||
fun_asr_nano,
|
fun_asr_nano,
|
||||||
device,
|
device,
|
||||||
dtype,
|
dtype,
|
||||||
// eos_token_id1: generation_config.eos_token_id[0] as u32,
|
|
||||||
// eos_token_id2: generation_config.eos_token_id[1] as u32,
|
|
||||||
generation_config,
|
generation_config,
|
||||||
model_name,
|
model_name,
|
||||||
})
|
})
|
||||||
@@ -165,58 +161,10 @@ impl GenerateModel for FunAsrNanoGenerateModel {
|
|||||||
top_k.into(),
|
top_k.into(),
|
||||||
seed,
|
seed,
|
||||||
max_tokens,
|
max_tokens,
|
||||||
|
false,
|
||||||
&self.device,
|
&self.device,
|
||||||
&self.model_name,
|
&self.model_name,
|
||||||
)?;
|
)?;
|
||||||
// let mut seq_len = input_ids.dim(1)?;
|
|
||||||
// let mut seqlen_offset = 0;
|
|
||||||
// let sample_len = mes.max_tokens.unwrap_or(1024);
|
|
||||||
// let stream = stream! {
|
|
||||||
// let mut error_tokens = Vec::new();
|
|
||||||
// let mut speech = Some(speech.to_dtype(self.dtype)?);
|
|
||||||
// let mut fbank_mask = Some(&fbank_mask);
|
|
||||||
// let mut input_ids = input_ids;
|
|
||||||
// for _ in 0..sample_len {
|
|
||||||
// let logits = self.fun_asr_nano.forward(
|
|
||||||
// &input_ids,
|
|
||||||
// speech.as_ref(),
|
|
||||||
// fbank_mask,
|
|
||||||
// seqlen_offset,
|
|
||||||
// )?;
|
|
||||||
// let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
|
||||||
// let next_token = logit_processor.sample(&logits)?;
|
|
||||||
// let mut decode_ids = Vec::new();
|
|
||||||
// if !error_tokens.is_empty() {
|
|
||||||
// decode_ids.extend_from_slice(&error_tokens);
|
|
||||||
// }
|
|
||||||
// decode_ids.push(next_token);
|
|
||||||
// let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
|
|
||||||
// if decoded_token.contains("�") {
|
|
||||||
// error_tokens.push(next_token);
|
|
||||||
// if error_tokens.len() > 3 {
|
|
||||||
// error_tokens.clear();
|
|
||||||
// }
|
|
||||||
// seqlen_offset += seq_len;
|
|
||||||
// seq_len = 1;
|
|
||||||
// input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
// speech = None;
|
|
||||||
// fbank_mask = None;
|
|
||||||
// continue;
|
|
||||||
// }
|
|
||||||
// error_tokens.clear();
|
|
||||||
// let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
|
||||||
// yield Ok(chunk);
|
|
||||||
// if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
|
||||||
// break;
|
|
||||||
// }
|
|
||||||
// seqlen_offset += seq_len;
|
|
||||||
// seq_len = 1;
|
|
||||||
// input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
// speech = None;
|
|
||||||
// fbank_mask = None;
|
|
||||||
// }
|
|
||||||
// self.fun_asr_nano.clear_kv_cache();
|
|
||||||
// };
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -611,7 +611,7 @@ impl FunAsrNanoModel {
|
|||||||
config.audio_adaptor_conf.n_layer,
|
config.audio_adaptor_conf.n_layer,
|
||||||
8,
|
8,
|
||||||
)?;
|
)?;
|
||||||
let llm = Qwen3Model::new(llm_cfg, vb.pp("llm"))?;
|
let llm = Qwen3Model::new(llm_cfg, vb.pp("llm"), vec![])?;
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
audio_encoder,
|
audio_encoder,
|
||||||
audio_adaptor,
|
audio_adaptor,
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
use crate::params::chat::ChatCompletionParameters;
|
use crate::params::chat::ChatCompletionParameters;
|
||||||
use anyhow::Result;
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{D, Device, Tensor};
|
use candle_core::{D, Device, Tensor};
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
@@ -92,6 +92,9 @@ impl FunAsrNanoProcessor {
|
|||||||
source_ids.extend_from_slice(&sub_token);
|
source_ids.extend_from_slice(&sub_token);
|
||||||
fbank_mask.extend_from_slice(&vec![0u32; sub_token.len()]);
|
fbank_mask.extend_from_slice(&vec![0u32; sub_token.len()]);
|
||||||
let audio_tensors = extract_audios(mes, &self.device, Some(self.fronted_conf.fs))?;
|
let audio_tensors = extract_audios(mes, &self.device, Some(self.fronted_conf.fs))?;
|
||||||
|
if audio_tensors.is_empty() {
|
||||||
|
return Err(anyhow!("FunASRNano need audio input"));
|
||||||
|
}
|
||||||
let audio = &audio_tensors[0];
|
let audio = &audio_tensors[0];
|
||||||
let (speech, speech_lengths) = self.extract_fbank(audio)?;
|
let (speech, speech_lengths) = self.extract_fbank(audio)?;
|
||||||
let olens = 1 + (speech_lengths - 3 + 2) / 2;
|
let olens = 1 + (speech_lengths - 3 + 2) / 2;
|
||||||
|
|||||||
@@ -1,11 +1,13 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::common::generate::get_logit_processor,
|
models::common::{
|
||||||
|
MultiModalData,
|
||||||
|
generate::{GenerationContext, generate_generic, generate_stream_generic},
|
||||||
|
},
|
||||||
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
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::async_stream::stream;
|
|
||||||
use rocket::futures::Stream;
|
use rocket::futures::Stream;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
@@ -17,10 +19,7 @@ use crate::{
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{find_type_files, get_device, get_dtype},
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
|
||||||
get_dtype,
|
|
||||||
},
|
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct GlmAsrNanoGenerateModel<'a> {
|
pub struct GlmAsrNanoGenerateModel<'a> {
|
||||||
@@ -30,9 +29,6 @@ pub struct GlmAsrNanoGenerateModel<'a> {
|
|||||||
glm_asr_nano: GlmAsrNanoModel,
|
glm_asr_nano: GlmAsrNanoModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
dtype: DType,
|
dtype: DType,
|
||||||
eos_token_id1: u32,
|
|
||||||
eos_token_id2: u32,
|
|
||||||
eos_token_id3: u32,
|
|
||||||
model_name: String,
|
model_name: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -48,7 +44,8 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
|
|||||||
let dtype = get_dtype(dtype, cfg_dtype);
|
let dtype = get_dtype(dtype, cfg_dtype);
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
||||||
let glm_asr_nano = GlmAsrNanoModel::new(vb, cfg)?;
|
let eos_ids = vec![59246u32, 59253, 59255];
|
||||||
|
let glm_asr_nano = GlmAsrNanoModel::new(vb, cfg, eos_ids)?;
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -61,9 +58,6 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
|
|||||||
glm_asr_nano,
|
glm_asr_nano,
|
||||||
device,
|
device,
|
||||||
dtype,
|
dtype,
|
||||||
eos_token_id1: 59246,
|
|
||||||
eos_token_id2: 59253,
|
|
||||||
eos_token_id3: 59255,
|
|
||||||
model_name,
|
model_name,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -72,46 +66,33 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
|
|||||||
impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
|
impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
|
||||||
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 mut input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
|
let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
|
||||||
let mut input_features = Some(input_features.to_dtype(self.dtype)?);
|
let input_features = input_features.to_dtype(self.dtype)?;
|
||||||
let mut audio_token_lengths = Some(audio_token_lengths);
|
let audio_token_lengths = Tensor::new(audio_token_lengths, &self.device)?;
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
|
||||||
let prompt_tokens = seq_len as u32;
|
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
let mut generate: Vec<u32> = Vec::new();
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
for _ in 0..sample_len {
|
let mut ctx = GenerationContext::new(
|
||||||
let logits = self.glm_asr_nano.forward(
|
mes.temperature,
|
||||||
input_features.as_ref(),
|
mes.top_p,
|
||||||
audio_token_lengths.as_ref(),
|
mes.top_k,
|
||||||
&input_ids,
|
seed,
|
||||||
seqlen_offset,
|
input_ids.dim(1)?,
|
||||||
)?;
|
sample_len,
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
self.device.clone(),
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
);
|
||||||
generate.push(next_token);
|
|
||||||
if next_token == self.eos_token_id1
|
let data_vec = vec![input_features.into(), audio_token_lengths.into()];
|
||||||
|| next_token == self.eos_token_id2
|
let data = MultiModalData::new(data_vec);
|
||||||
|| next_token == self.eos_token_id3
|
generate_generic(
|
||||||
{
|
&mut self.glm_asr_nano,
|
||||||
break;
|
&self.tokenizer,
|
||||||
}
|
input_ids,
|
||||||
seqlen_offset += seq_len;
|
data,
|
||||||
seq_len = 1;
|
&mut ctx,
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
&self.model_name,
|
||||||
input_features = None;
|
)
|
||||||
audio_token_lengths = None;
|
|
||||||
}
|
|
||||||
let num_token = generate.len() as u32;
|
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
|
||||||
self.glm_asr_nano.clear_kv_cache();
|
|
||||||
let response =
|
|
||||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
|
||||||
Ok(response)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
@@ -126,58 +107,29 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
|
|||||||
>,
|
>,
|
||||||
> {
|
> {
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
|
||||||
let render_text = self.chat_template.apply_chat_template(&mes)?;
|
let render_text = 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 mut seq_len = input_ids.dim(1)?;
|
let audio_token_lengths = Tensor::new(audio_token_lengths, &self.device)?;
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
let stream = stream! {
|
let data_vec = vec![input_features.into(), audio_token_lengths.into()];
|
||||||
let mut error_tokens = Vec::new();
|
let data = MultiModalData::new(data_vec);
|
||||||
let mut input_features = Some(input_features.to_dtype(self.dtype)?);
|
let stream = generate_stream_generic(
|
||||||
let mut audio_token_lengths = Some(audio_token_lengths);
|
&mut self.glm_asr_nano,
|
||||||
let mut input_ids = input_ids;
|
&self.tokenizer,
|
||||||
for _ in 0..sample_len {
|
input_ids,
|
||||||
let logits =
|
data,
|
||||||
self.glm_asr_nano
|
mes.temperature,
|
||||||
.forward(input_features.as_ref(), audio_token_lengths.as_ref(), &input_ids, seqlen_offset)?;
|
mes.top_p,
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
None,
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
seed,
|
||||||
let mut decode_ids = Vec::new();
|
sample_len,
|
||||||
if !error_tokens.is_empty() {
|
false,
|
||||||
decode_ids.extend_from_slice(&error_tokens);
|
&self.device,
|
||||||
}
|
&self.model_name,
|
||||||
decode_ids.push(next_token);
|
)?;
|
||||||
let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
|
|
||||||
if decoded_token.contains("�") {
|
|
||||||
error_tokens.push(next_token);
|
|
||||||
if error_tokens.len() > 3 {
|
|
||||||
error_tokens.clear();
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
input_features = None;
|
|
||||||
audio_token_lengths = None;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
error_tokens.clear();
|
|
||||||
let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
|
||||||
yield Ok(chunk);
|
|
||||||
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 || next_token == self.eos_token_id3{
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
input_features = None;
|
|
||||||
audio_token_lengths = None;
|
|
||||||
}
|
|
||||||
self.glm_asr_nano.clear_kv_cache();
|
|
||||||
};
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,9 +4,12 @@ use candle_nn::{Conv1d, LayerNorm, Linear, Module, VarBuilder, linear, linear_no
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::modules::{
|
common::{
|
||||||
|
InferenceModel,
|
||||||
|
modules::{
|
||||||
LlamaForCausalLM, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm,
|
LlamaForCausalLM, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm,
|
||||||
},
|
},
|
||||||
|
},
|
||||||
glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig},
|
glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig},
|
||||||
},
|
},
|
||||||
position_embed::rope::{RoPE, glm_asr_apply_rotary_pos_emb},
|
position_embed::rope::{RoPE, glm_asr_apply_rotary_pos_emb},
|
||||||
@@ -233,10 +236,11 @@ pub struct GlmAsrNanoModel {
|
|||||||
audio_tower: GlmAsrEncoder,
|
audio_tower: GlmAsrEncoder,
|
||||||
multi_modal_projector: TwoLinearMLP,
|
multi_modal_projector: TwoLinearMLP,
|
||||||
language_model: LlamaForCausalLM,
|
language_model: LlamaForCausalLM,
|
||||||
|
stop_token_ids: Vec<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl GlmAsrNanoModel {
|
impl GlmAsrNanoModel {
|
||||||
pub fn new(vb: VarBuilder, config: GlmAsrNanoConfig) -> Result<Self> {
|
pub fn new(vb: VarBuilder, config: GlmAsrNanoConfig, eos_ids: Vec<u32>) -> Result<Self> {
|
||||||
let audio_tower = GlmAsrEncoder::new(vb.pp("audio_tower"), &config.audio_config)?;
|
let audio_tower = GlmAsrEncoder::new(vb.pp("audio_tower"), &config.audio_config)?;
|
||||||
let multi_modal_projector = TwoLinearMLP::new(
|
let multi_modal_projector = TwoLinearMLP::new(
|
||||||
vb.pp("multi_modal_projector"),
|
vb.pp("multi_modal_projector"),
|
||||||
@@ -273,6 +277,7 @@ impl GlmAsrNanoModel {
|
|||||||
audio_tower,
|
audio_tower,
|
||||||
multi_modal_projector,
|
multi_modal_projector,
|
||||||
language_model,
|
language_model,
|
||||||
|
stop_token_ids: eos_ids,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -300,7 +305,7 @@ impl GlmAsrNanoModel {
|
|||||||
pub fn forward(
|
pub fn forward(
|
||||||
&mut self,
|
&mut self,
|
||||||
input_features: Option<&Tensor>,
|
input_features: Option<&Tensor>,
|
||||||
audio_token_lengths: Option<&Vec<u32>>,
|
audio_token_lengths: Option<&Tensor>,
|
||||||
input_ids: &Tensor,
|
input_ids: &Tensor,
|
||||||
seqlen_offset: usize,
|
seqlen_offset: usize,
|
||||||
) -> Result<Tensor> {
|
) -> Result<Tensor> {
|
||||||
@@ -308,8 +313,9 @@ impl GlmAsrNanoModel {
|
|||||||
if let Some(input_features) = input_features
|
if let Some(input_features) = input_features
|
||||||
&& let Some(audio_token_len) = audio_token_lengths
|
&& let Some(audio_token_len) = audio_token_lengths
|
||||||
{
|
{
|
||||||
|
let audio_token_len = audio_token_len.to_vec1::<u32>()?;
|
||||||
let audio_token_mask = get_equal_mask(input_ids, self.config.audio_token_id)?;
|
let audio_token_mask = get_equal_mask(input_ids, self.config.audio_token_id)?;
|
||||||
let audio_embeds = self.get_audio_features(input_features, audio_token_len)?;
|
let audio_embeds = self.get_audio_features(input_features, &audio_token_len)?;
|
||||||
inputs_embeds = masked_scatter_dim0(&inputs_embeds, &audio_embeds, &audio_token_mask)?;
|
inputs_embeds = masked_scatter_dim0(&inputs_embeds, &audio_embeds, &audio_token_mask)?;
|
||||||
}
|
}
|
||||||
let logits = self.language_model.forward(&inputs_embeds, seqlen_offset)?;
|
let logits = self.language_model.forward(&inputs_embeds, seqlen_offset)?;
|
||||||
@@ -319,3 +325,38 @@ impl GlmAsrNanoModel {
|
|||||||
self.language_model.clear_kv_cache();
|
self.language_model.clear_kv_cache();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl InferenceModel for GlmAsrNanoModel {
|
||||||
|
fn forward_initial(
|
||||||
|
&mut self,
|
||||||
|
input_ids: &Tensor,
|
||||||
|
seqlen_offset: usize,
|
||||||
|
data: crate::models::common::MultiModalData,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
if data.data_vec.len() != 2 {
|
||||||
|
return Err(anyhow::anyhow!(
|
||||||
|
"GlmAsrNano process data error, must have input_features, audio_token_lengths"
|
||||||
|
));
|
||||||
|
}
|
||||||
|
let input_features = &data.data_vec[0];
|
||||||
|
let audio_token_lengths = &data.data_vec[1];
|
||||||
|
self.forward(
|
||||||
|
input_features.as_ref(),
|
||||||
|
audio_token_lengths.as_ref(),
|
||||||
|
input_ids,
|
||||||
|
seqlen_offset,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
|
self.forward(None, None, input_ids, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clear_cache(&mut self) {
|
||||||
|
self.clear_kv_cache();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn stop_token_ids(&self) -> Vec<u32> {
|
||||||
|
self.stop_token_ids.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -206,6 +206,9 @@ impl GlmAsrNanoProcessor {
|
|||||||
render_text: &str,
|
render_text: &str,
|
||||||
) -> Result<(Tensor, Vec<u32>, String)> {
|
) -> Result<(Tensor, Vec<u32>, String)> {
|
||||||
let audio_tensors = extract_audios(mes, &self.device, Some(self.sampling_rate))?;
|
let audio_tensors = extract_audios(mes, &self.device, Some(self.sampling_rate))?;
|
||||||
|
if audio_tensors.is_empty() {
|
||||||
|
return Err(anyhow::anyhow!("GlmASRNano need audio input"));
|
||||||
|
}
|
||||||
let (input_features, input_features_mask, per_sample_windows) =
|
let (input_features, input_features_mask, per_sample_windows) =
|
||||||
self.process_audio(audio_tensors)?;
|
self.process_audio(audio_tensors)?;
|
||||||
let audio_lengths = input_features_mask.sum(D::Minus1)?;
|
let audio_lengths = input_features_mask.sum(D::Minus1)?;
|
||||||
|
|||||||
+54
-117
@@ -1,12 +1,14 @@
|
|||||||
//! GLM-OCR Inference and Generation
|
//! GLM-OCR Inference and Generation
|
||||||
use crate::{
|
use crate::{
|
||||||
models::common::generate::get_logit_processor,
|
models::common::{
|
||||||
|
MultiModalData,
|
||||||
|
generate::{GenerationContext, generate_generic, generate_stream_generic},
|
||||||
|
},
|
||||||
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{DType, Device, IndexOp, Tensor};
|
use candle_core::{DType, Device};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use rocket::async_stream::stream;
|
|
||||||
use rocket::futures::Stream;
|
use rocket::futures::Stream;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
@@ -21,8 +23,7 @@ use crate::{
|
|||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{
|
||||||
build_completion_chunk_response, build_completion_response, extract_user_text,
|
extract_user_text, find_type_files, get_device, get_dtype, img_utils::extract_image_url,
|
||||||
find_type_files, get_device, get_dtype, img_utils::extract_image_url,
|
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -32,7 +33,6 @@ pub struct GlmOcrGenerateModel {
|
|||||||
processor: GlmOcrProcessor,
|
processor: GlmOcrProcessor,
|
||||||
model: GlmOcrModel,
|
model: GlmOcrModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
eos_token_ids: Vec<u32>,
|
|
||||||
model_name: String,
|
model_name: String,
|
||||||
image_token_id: u32,
|
image_token_id: u32,
|
||||||
image_start_token_id: u32,
|
image_start_token_id: u32,
|
||||||
@@ -54,10 +54,10 @@ impl GlmOcrGenerateModel {
|
|||||||
let processor = GlmOcrProcessor::new(path, &device, dtype)?;
|
let processor = GlmOcrProcessor::new(path, &device, dtype)?;
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
||||||
let model = GlmOcrModel::new(vb, cfg.clone())?;
|
|
||||||
let generation_config_path = path.to_string() + "/generation_config.json";
|
let generation_config_path = path.to_string() + "/generation_config.json";
|
||||||
let generation_config: GlmOcrGenerationConfig =
|
let generation_config: GlmOcrGenerationConfig =
|
||||||
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
||||||
|
let model = GlmOcrModel::new(vb, cfg.clone(), generation_config.eos_token_id.clone())?;
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -69,7 +69,6 @@ impl GlmOcrGenerateModel {
|
|||||||
processor,
|
processor,
|
||||||
model,
|
model,
|
||||||
device,
|
device,
|
||||||
eos_token_ids: generation_config.eos_token_id.clone(),
|
|
||||||
model_name,
|
model_name,
|
||||||
image_token_id: cfg.image_token_id,
|
image_token_id: cfg.image_token_id,
|
||||||
image_start_token_id: cfg.image_start_token_id,
|
image_start_token_id: cfg.image_start_token_id,
|
||||||
@@ -84,8 +83,6 @@ impl GlmOcrGenerateModel {
|
|||||||
impl GenerateModel for GlmOcrGenerateModel {
|
impl GenerateModel for GlmOcrGenerateModel {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
|
||||||
|
|
||||||
// Extract image path and prompt from messages
|
// Extract image path and prompt from messages
|
||||||
let image_urls = extract_image_url(&mes);
|
let image_urls = extract_image_url(&mes);
|
||||||
let image_path = image_urls
|
let image_path = image_urls
|
||||||
@@ -110,55 +107,32 @@ impl GenerateModel for GlmOcrGenerateModel {
|
|||||||
self.spatial_merge_size,
|
self.spatial_merge_size,
|
||||||
)?;
|
)?;
|
||||||
|
|
||||||
let mut input_ids = processed.input_ids;
|
let input_ids = processed.input_ids;
|
||||||
let pixel_values = Some(processed.pixel_values);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
let image_grid_thw = Some(processed.grid_thw);
|
let mut ctx = GenerationContext::new(
|
||||||
let image_mask = Some(processed.image_mask);
|
mes.temperature,
|
||||||
let mut seqlen_offset = 0;
|
mes.top_p,
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
mes.top_k,
|
||||||
let prompt_tokens = seq_len as u32;
|
seed,
|
||||||
let mut generate = Vec::new();
|
input_ids.dim(1)?,
|
||||||
let sample_len = mes.max_tokens.unwrap_or(512);
|
sample_len,
|
||||||
|
self.device.clone(),
|
||||||
|
);
|
||||||
|
|
||||||
for _ in 0..sample_len {
|
let data_vec = vec![
|
||||||
let is_first_pass = seqlen_offset == 0;
|
processed.pixel_values.into(),
|
||||||
let logits = self.model.forward(
|
processed.grid_thw.into(),
|
||||||
&input_ids,
|
processed.image_mask.into(),
|
||||||
if is_first_pass {
|
];
|
||||||
pixel_values.as_ref()
|
let data = MultiModalData::new(data_vec);
|
||||||
} else {
|
generate_generic(
|
||||||
None
|
&mut self.model,
|
||||||
},
|
&self.tokenizer,
|
||||||
if is_first_pass {
|
input_ids,
|
||||||
image_grid_thw.as_ref()
|
data,
|
||||||
} else {
|
&mut ctx,
|
||||||
None
|
&self.model_name,
|
||||||
},
|
)
|
||||||
if is_first_pass {
|
|
||||||
image_mask.as_ref()
|
|
||||||
} else {
|
|
||||||
None
|
|
||||||
},
|
|
||||||
seqlen_offset,
|
|
||||||
)?;
|
|
||||||
let logits = logits.i((0, seq_len - 1, ..))?.to_dtype(DType::F32)?;
|
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
|
||||||
|
|
||||||
generate.push(next_token);
|
|
||||||
if self.eos_token_ids.contains(&next_token) {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
}
|
|
||||||
|
|
||||||
self.model.clear_kv_cache();
|
|
||||||
let num_token = generate.len() as u32;
|
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
|
||||||
let response =
|
|
||||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
|
||||||
Ok(response)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
@@ -173,8 +147,6 @@ impl GenerateModel for GlmOcrGenerateModel {
|
|||||||
>,
|
>,
|
||||||
> {
|
> {
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
|
||||||
|
|
||||||
// Extract image path and prompt from messages
|
// Extract image path and prompt from messages
|
||||||
let image_urls = extract_image_url(&mes);
|
let image_urls = extract_image_url(&mes);
|
||||||
let image_path = image_urls
|
let image_path = image_urls
|
||||||
@@ -199,63 +171,28 @@ impl GenerateModel for GlmOcrGenerateModel {
|
|||||||
self.spatial_merge_size,
|
self.spatial_merge_size,
|
||||||
)?;
|
)?;
|
||||||
|
|
||||||
let mut input_ids = processed.input_ids;
|
let input_ids = processed.input_ids;
|
||||||
let pixel_values = Some(processed.pixel_values);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
let image_grid_thw = Some(processed.grid_thw);
|
let data_vec = vec![
|
||||||
let image_mask = Some(processed.image_mask);
|
processed.pixel_values.into(),
|
||||||
let mut seqlen_offset = 0;
|
processed.grid_thw.into(),
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
processed.image_mask.into(),
|
||||||
let sample_len = mes.max_tokens.unwrap_or(512);
|
];
|
||||||
|
let data = MultiModalData::new(data_vec);
|
||||||
let stream = stream! {
|
let stream = generate_stream_generic(
|
||||||
let mut generated: Vec<u32> = Vec::new();
|
&mut self.model,
|
||||||
let mut error_tokens = Vec::new();
|
&self.tokenizer,
|
||||||
for _ in 0..sample_len {
|
input_ids,
|
||||||
let is_first_pass = seqlen_offset == 0;
|
data,
|
||||||
let logits = self.model.forward(
|
mes.temperature,
|
||||||
&input_ids,
|
mes.top_p,
|
||||||
if is_first_pass { pixel_values.as_ref() } else { None },
|
None,
|
||||||
if is_first_pass { image_grid_thw.as_ref() } else { None },
|
seed,
|
||||||
if is_first_pass { image_mask.as_ref() } else { None },
|
sample_len,
|
||||||
seqlen_offset,
|
false,
|
||||||
).map_err(|e| anyhow!(format!("forward error: {e}")))?;
|
&self.device,
|
||||||
let logits = logits.i((0, seq_len - 1, ..)).map_err(|e| anyhow!(format!("index error: {e}")))?.to_dtype(DType::F32).map_err(|e| anyhow!(format!("dtype error: {e}")))?;
|
&self.model_name,
|
||||||
|
)?;
|
||||||
let next_token = logit_processor.sample(&logits).map_err(|e| anyhow!(format!("sample error: {e}")))?;
|
|
||||||
generated.push(next_token);
|
|
||||||
|
|
||||||
let mut decode_ids = Vec::new();
|
|
||||||
if !error_tokens.is_empty() {
|
|
||||||
decode_ids.extend_from_slice(&error_tokens);
|
|
||||||
}
|
|
||||||
decode_ids.push(next_token);
|
|
||||||
|
|
||||||
let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("decode error: {e}")))?;
|
|
||||||
if decoded_token.contains("�") {
|
|
||||||
error_tokens.push(next_token);
|
|
||||||
if error_tokens.len() > 3 {
|
|
||||||
error_tokens.clear();
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device).map_err(|e| anyhow!(format!("tensor error: {e}")))?;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
error_tokens.clear();
|
|
||||||
|
|
||||||
let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
|
||||||
yield Ok(chunk);
|
|
||||||
|
|
||||||
if self.eos_token_ids.contains(&next_token) {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device).map_err(|e| anyhow!(format!("tensor error: {e}")))?;
|
|
||||||
}
|
|
||||||
self.model.clear_kv_cache();
|
|
||||||
};
|
|
||||||
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ use candle_nn::{
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::modules::GateUpDownMLP,
|
common::{InferenceModel, modules::GateUpDownMLP},
|
||||||
glm_ocr::config::{GlmOcrConfig, GlmOcrTextConfig, GlmOcrVisionConfig},
|
glm_ocr::config::{GlmOcrConfig, GlmOcrTextConfig, GlmOcrVisionConfig},
|
||||||
},
|
},
|
||||||
position_embed::rope::{apply_rotary_pos_emb_vision, glm_ocr_apply_rotary_pos_emb},
|
position_embed::rope::{apply_rotary_pos_emb_vision, glm_ocr_apply_rotary_pos_emb},
|
||||||
@@ -1256,7 +1256,8 @@ impl GlmOcrTextModel {
|
|||||||
}
|
}
|
||||||
|
|
||||||
hidden_states = self.norm.forward(&hidden_states)?;
|
hidden_states = self.norm.forward(&hidden_states)?;
|
||||||
let logits = self.lm_head.forward(&hidden_states)?;
|
let last = hidden_states.narrow(1, seq_len - 1, 1)?;
|
||||||
|
let logits = self.lm_head.forward(&last)?;
|
||||||
|
|
||||||
Ok(logits)
|
Ok(logits)
|
||||||
}
|
}
|
||||||
@@ -1271,10 +1272,11 @@ impl GlmOcrTextModel {
|
|||||||
pub struct GlmOcrModel {
|
pub struct GlmOcrModel {
|
||||||
vision_encoder: GlmOcrVisionModel,
|
vision_encoder: GlmOcrVisionModel,
|
||||||
language_model: GlmOcrTextModel,
|
language_model: GlmOcrTextModel,
|
||||||
|
stop_token_ids: Vec<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl GlmOcrModel {
|
impl GlmOcrModel {
|
||||||
pub fn new(vb: VarBuilder, config: GlmOcrConfig) -> Result<Self> {
|
pub fn new(vb: VarBuilder, config: GlmOcrConfig, eos_ids: Vec<u32>) -> Result<Self> {
|
||||||
let vision_encoder =
|
let vision_encoder =
|
||||||
GlmOcrVisionModel::new(vb.pp("model").pp("visual"), &config.vision_config)?;
|
GlmOcrVisionModel::new(vb.pp("model").pp("visual"), &config.vision_config)?;
|
||||||
let language_model = GlmOcrTextModel::new(
|
let language_model = GlmOcrTextModel::new(
|
||||||
@@ -1286,6 +1288,7 @@ impl GlmOcrModel {
|
|||||||
Ok(Self {
|
Ok(Self {
|
||||||
vision_encoder,
|
vision_encoder,
|
||||||
language_model,
|
language_model,
|
||||||
|
stop_token_ids: eos_ids,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1331,3 +1334,40 @@ impl GlmOcrModel {
|
|||||||
self.language_model.clear_kv_cache();
|
self.language_model.clear_kv_cache();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl InferenceModel for GlmOcrModel {
|
||||||
|
fn forward_initial(
|
||||||
|
&mut self,
|
||||||
|
input_ids: &Tensor,
|
||||||
|
seqlen_offset: usize,
|
||||||
|
data: crate::models::common::MultiModalData,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
if data.data_vec.len() != 3 {
|
||||||
|
return Err(anyhow::anyhow!(
|
||||||
|
"GlmOcr process data error, must have pixel_values, image_grid_thw, image_mask"
|
||||||
|
));
|
||||||
|
}
|
||||||
|
let pixel_values = &data.data_vec[0];
|
||||||
|
let image_grid_thw = &data.data_vec[1];
|
||||||
|
let image_mask = &data.data_vec[2];
|
||||||
|
self.forward(
|
||||||
|
input_ids,
|
||||||
|
pixel_values.as_ref(),
|
||||||
|
image_grid_thw.as_ref(),
|
||||||
|
image_mask.as_ref(),
|
||||||
|
seqlen_offset,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
|
self.forward(input_ids, None, None, None, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clear_cache(&mut self) {
|
||||||
|
self.clear_kv_cache();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn stop_token_ids(&self) -> Vec<u32> {
|
||||||
|
self.stop_token_ids.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -85,7 +85,7 @@ pub struct HunyuanOCRGenerationConfig {
|
|||||||
pub bos_token_id: usize,
|
pub bos_token_id: usize,
|
||||||
pub pad_token_id: usize,
|
pub pad_token_id: usize,
|
||||||
pub do_sample: bool,
|
pub do_sample: bool,
|
||||||
pub eos_token_id: Vec<usize>,
|
pub eos_token_id: Vec<u32>,
|
||||||
pub top_p: f32,
|
pub top_p: f32,
|
||||||
pub top_k: usize,
|
pub top_k: usize,
|
||||||
pub temperature: f32,
|
pub temperature: f32,
|
||||||
|
|||||||
@@ -1,11 +1,13 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::common::generate::get_logit_processor,
|
models::common::{
|
||||||
|
MultiModalData,
|
||||||
|
generate::{GenerationContext, generate_generic, generate_stream_generic},
|
||||||
|
},
|
||||||
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use rocket::async_stream::stream;
|
|
||||||
use rocket::futures::Stream;
|
use rocket::futures::Stream;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
@@ -19,10 +21,7 @@ use crate::{
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{find_type_files, get_device, get_dtype},
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
|
||||||
get_dtype,
|
|
||||||
},
|
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct HunyuanOCRGenerateModel<'a> {
|
pub struct HunyuanOCRGenerateModel<'a> {
|
||||||
@@ -31,8 +30,6 @@ pub struct HunyuanOCRGenerateModel<'a> {
|
|||||||
pre_processor: HunyuanVLProcessor,
|
pre_processor: HunyuanVLProcessor,
|
||||||
hunyuan_vl: HunyuanVLModel,
|
hunyuan_vl: HunyuanVLModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
eos_token_id1: u32,
|
|
||||||
eos_token_id2: u32,
|
|
||||||
generation_config: HunyuanOCRGenerationConfig,
|
generation_config: HunyuanOCRGenerationConfig,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
}
|
}
|
||||||
@@ -49,10 +46,12 @@ impl<'a> HunyuanOCRGenerateModel<'a> {
|
|||||||
let pre_processor = HunyuanVLProcessor::new(path, &device, dtype)?;
|
let pre_processor = HunyuanVLProcessor::new(path, &device, dtype)?;
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
||||||
let hunyuan_vl = HunyuanVLModel::new(vb, cfg.clone())?;
|
|
||||||
let generation_config_path = path.to_string() + "/generation_config.json";
|
let generation_config_path = path.to_string() + "/generation_config.json";
|
||||||
let generation_config: HunyuanOCRGenerationConfig =
|
let generation_config: HunyuanOCRGenerationConfig =
|
||||||
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
||||||
|
let hunyuan_vl =
|
||||||
|
HunyuanVLModel::new(vb, cfg.clone(), generation_config.eos_token_id.clone())?;
|
||||||
|
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -64,8 +63,6 @@ impl<'a> HunyuanOCRGenerateModel<'a> {
|
|||||||
pre_processor,
|
pre_processor,
|
||||||
hunyuan_vl,
|
hunyuan_vl,
|
||||||
device,
|
device,
|
||||||
eos_token_id1: generation_config.eos_token_id[0] as u32,
|
|
||||||
eos_token_id2: generation_config.eos_token_id[1] as u32,
|
|
||||||
generation_config,
|
generation_config,
|
||||||
model_name,
|
model_name,
|
||||||
})
|
})
|
||||||
@@ -80,51 +77,37 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
|
|||||||
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
||||||
let top_k = self.generation_config.top_k;
|
let top_k = self.generation_config.top_k;
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let mut logit_processor =
|
|
||||||
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
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 mut input_ids = data.input_ids;
|
|
||||||
let mut position_ids = Some(&data.position_ids);
|
|
||||||
let mut image_mask = Some(&data.image_mask);
|
|
||||||
let mut pixel_values = data.pixel_values;
|
|
||||||
let mut image_grid_thw = data.image_grid_thw;
|
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
|
||||||
let prompt_tokens = seq_len as u32;
|
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
let mut generate: Vec<u32> = Vec::new();
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
for _ in 0..sample_len {
|
let input_ids = data.input_ids;
|
||||||
let logits = self.hunyuan_vl.forward(
|
let mut ctx = GenerationContext::new(
|
||||||
&input_ids,
|
temperature.into(),
|
||||||
pixel_values.as_ref(),
|
top_p.into(),
|
||||||
image_grid_thw.as_ref(),
|
top_k.into(),
|
||||||
image_mask,
|
seed,
|
||||||
position_ids,
|
input_ids.dim(1)?,
|
||||||
seqlen_offset,
|
sample_len,
|
||||||
)?;
|
self.device.clone(),
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
);
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
|
||||||
generate.push(next_token);
|
let data_vec = vec![
|
||||||
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
data.pixel_values,
|
||||||
break;
|
data.image_grid_thw,
|
||||||
}
|
data.image_mask.into(),
|
||||||
seqlen_offset += seq_len;
|
data.position_ids.into(),
|
||||||
seq_len = 1;
|
];
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
let data = MultiModalData::new(data_vec);
|
||||||
position_ids = None;
|
generate_generic(
|
||||||
image_mask = None;
|
&mut self.hunyuan_vl,
|
||||||
pixel_values = None;
|
&self.tokenizer,
|
||||||
image_grid_thw = None;
|
input_ids,
|
||||||
}
|
data,
|
||||||
let num_token = generate.len() as u32;
|
&mut ctx,
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
&self.model_name,
|
||||||
self.hunyuan_vl.clear_kv_cache();
|
)
|
||||||
let response =
|
|
||||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
|
||||||
Ok(response)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
@@ -144,70 +127,34 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
|
|||||||
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
||||||
let top_k = self.generation_config.top_k;
|
let top_k = self.generation_config.top_k;
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let mut logit_processor =
|
|
||||||
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
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 mut seqlen_offset = 0;
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
let stream = stream! {
|
let input_ids = data.input_ids;
|
||||||
let mut error_tokens = Vec::new();
|
let data_vec = vec![
|
||||||
let mut input_ids = data.input_ids;
|
data.pixel_values,
|
||||||
let mut position_ids = Some(&data.position_ids);
|
data.image_grid_thw,
|
||||||
let mut image_mask = Some(&data.image_mask);
|
data.image_mask.into(),
|
||||||
let mut pixel_values = data.pixel_values;
|
data.position_ids.into(),
|
||||||
let mut image_grid_thw = data.image_grid_thw;
|
];
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
let data = MultiModalData::new(data_vec);
|
||||||
for _ in 0..sample_len {
|
let stream = generate_stream_generic(
|
||||||
let logits = self.hunyuan_vl.forward(
|
&mut self.hunyuan_vl,
|
||||||
&input_ids,
|
&self.tokenizer,
|
||||||
pixel_values.as_ref(),
|
input_ids,
|
||||||
image_grid_thw.as_ref(),
|
data,
|
||||||
image_mask,
|
temperature.into(),
|
||||||
position_ids,
|
top_p.into(),
|
||||||
seqlen_offset,
|
top_k.into(),
|
||||||
|
seed,
|
||||||
|
sample_len,
|
||||||
|
false,
|
||||||
|
&self.device,
|
||||||
|
&self.model_name,
|
||||||
)?;
|
)?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
|
||||||
let mut decode_ids = Vec::new();
|
|
||||||
if !error_tokens.is_empty() {
|
|
||||||
decode_ids.extend_from_slice(&error_tokens);
|
|
||||||
}
|
|
||||||
decode_ids.push(next_token);
|
|
||||||
let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
|
|
||||||
if decoded_token.contains("�") {
|
|
||||||
error_tokens.push(next_token);
|
|
||||||
if error_tokens.len() > 3 {
|
|
||||||
error_tokens.clear();
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
position_ids = None;
|
|
||||||
image_mask = None;
|
|
||||||
pixel_values = None;
|
|
||||||
image_grid_thw = None;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
error_tokens.clear();
|
|
||||||
let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
|
||||||
yield Ok(chunk);
|
|
||||||
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
position_ids = None;
|
|
||||||
image_mask = None;
|
|
||||||
pixel_values = None;
|
|
||||||
image_grid_thw = None;
|
|
||||||
}
|
|
||||||
self.hunyuan_vl.clear_kv_cache();
|
|
||||||
};
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -7,14 +7,19 @@ use candle_nn::{
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::modules::{
|
common::{
|
||||||
|
InferenceModel,
|
||||||
|
modules::{
|
||||||
GateUpDownMLP, NaiveAttnTwoLinearMLPBlock, eager_attention_forward, get_conv2d,
|
GateUpDownMLP, NaiveAttnTwoLinearMLPBlock, eager_attention_forward, get_conv2d,
|
||||||
},
|
},
|
||||||
|
},
|
||||||
hunyuan_ocr::config::{HunYuanVLConfig, HunYuanVLVisionConfig},
|
hunyuan_ocr::config::{HunYuanVLConfig, HunYuanVLVisionConfig},
|
||||||
},
|
},
|
||||||
position_embed::rope::{RoPE, apply_rotary_pos_emb, get_xd_cos_sin},
|
position_embed::rope::{RoPE, apply_rotary_pos_emb, get_xd_cos_sin},
|
||||||
utils::interpolate::interpolate_bilinear,
|
utils::{
|
||||||
utils::tensor_utils::{masked_scatter_dim0, prepare_causal_attention_mask, split_tensor},
|
interpolate::interpolate_bilinear,
|
||||||
|
tensor_utils::{masked_scatter_dim0, prepare_causal_attention_mask, split_tensor},
|
||||||
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct HunYuanVisionPatchEmbed {
|
pub struct HunYuanVisionPatchEmbed {
|
||||||
@@ -538,10 +543,11 @@ pub struct HunyuanVLModel {
|
|||||||
vit: HunYuanVisionTransformer,
|
vit: HunYuanVisionTransformer,
|
||||||
model: HunYuanVLTextModel,
|
model: HunYuanVLTextModel,
|
||||||
lm_head: Linear,
|
lm_head: Linear,
|
||||||
|
stop_token_ids: Vec<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl HunyuanVLModel {
|
impl HunyuanVLModel {
|
||||||
pub fn new(vb: VarBuilder, config: HunYuanVLConfig) -> Result<Self> {
|
pub fn new(vb: VarBuilder, config: HunYuanVLConfig, eos_ids: Vec<u32>) -> Result<Self> {
|
||||||
let vit = HunYuanVisionTransformer::new(vb.pp("vit"), &config.vision_config)?;
|
let vit = HunYuanVisionTransformer::new(vb.pp("vit"), &config.vision_config)?;
|
||||||
let model = HunYuanVLTextModel::new(vb.pp("model"), &config)?;
|
let model = HunYuanVLTextModel::new(vb.pp("model"), &config)?;
|
||||||
let lm_head = Linear::new(model.embed_tokens.embeddings().clone(), None);
|
let lm_head = Linear::new(model.embed_tokens.embeddings().clone(), None);
|
||||||
@@ -550,6 +556,7 @@ impl HunyuanVLModel {
|
|||||||
vit,
|
vit,
|
||||||
model,
|
model,
|
||||||
lm_head,
|
lm_head,
|
||||||
|
stop_token_ids: eos_ids,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
pub fn forward(
|
pub fn forward(
|
||||||
@@ -582,3 +589,42 @@ impl HunyuanVLModel {
|
|||||||
self.model.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl InferenceModel for HunyuanVLModel {
|
||||||
|
fn forward_initial(
|
||||||
|
&mut self,
|
||||||
|
input_ids: &Tensor,
|
||||||
|
seqlen_offset: usize,
|
||||||
|
data: crate::models::common::MultiModalData,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
if data.data_vec.len() != 4 {
|
||||||
|
return Err(anyhow::anyhow!(
|
||||||
|
"HunyuanVL process data error, must have pixel_values, image_grid_thw, image_mask, position_ids"
|
||||||
|
));
|
||||||
|
}
|
||||||
|
let pixel_values = &data.data_vec[0];
|
||||||
|
let image_grid_thw = &data.data_vec[1];
|
||||||
|
let image_mask = &data.data_vec[2];
|
||||||
|
let position_ids = &data.data_vec[3];
|
||||||
|
self.forward(
|
||||||
|
input_ids,
|
||||||
|
pixel_values.as_ref(),
|
||||||
|
image_grid_thw.as_ref(),
|
||||||
|
image_mask.as_ref(),
|
||||||
|
position_ids.as_ref(),
|
||||||
|
seqlen_offset,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
|
self.forward(input_ids, None, None, None, None, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clear_cache(&mut self) {
|
||||||
|
self.clear_kv_cache();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn stop_token_ids(&self) -> Vec<u32> {
|
||||||
|
self.stop_token_ids.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+40
-79
@@ -1,6 +1,8 @@
|
|||||||
use crate::models::common::generate::get_logit_processor;
|
use crate::models::common::MultiModalData;
|
||||||
|
use crate::models::common::generate::{
|
||||||
|
GenerationContext, generate_generic, generate_stream_generic,
|
||||||
|
};
|
||||||
use crate::params::chat::{ChatCompletionParameters, ChatCompletionResponse};
|
use crate::params::chat::{ChatCompletionParameters, ChatCompletionResponse};
|
||||||
use crate::utils::build_completion_chunk_response;
|
|
||||||
use crate::{
|
use crate::{
|
||||||
chat_template::ChatTemplate,
|
chat_template::ChatTemplate,
|
||||||
models::{
|
models::{
|
||||||
@@ -11,19 +13,17 @@ use crate::{
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{build_completion_response, find_type_files, get_device, get_dtype},
|
utils::{find_type_files, get_device, get_dtype},
|
||||||
};
|
};
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use rocket::async_stream::stream;
|
|
||||||
|
|
||||||
pub struct Lfm2GenerateModel<'a> {
|
pub struct Lfm2GenerateModel<'a> {
|
||||||
chat_template: ChatTemplate<'a>,
|
chat_template: ChatTemplate<'a>,
|
||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
model: Lfm2Model,
|
model: Lfm2Model,
|
||||||
eos_token_id: u32,
|
|
||||||
model_name: String,
|
model_name: String,
|
||||||
}
|
}
|
||||||
impl<'a> Lfm2GenerateModel<'a> {
|
impl<'a> Lfm2GenerateModel<'a> {
|
||||||
@@ -45,8 +45,8 @@ impl<'a> Lfm2GenerateModel<'a> {
|
|||||||
};
|
};
|
||||||
let dtype = get_dtype(dtype, &cfg_dtype);
|
let dtype = get_dtype(dtype, &cfg_dtype);
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_path, dtype, &device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_path, dtype, &device)? };
|
||||||
let model = Lfm2Model::new(vb, &cfg)?;
|
let eos_ids = vec![gen_cfg.eos_token_id];
|
||||||
let eos_token_id = gen_cfg.eos_token_id;
|
let model = Lfm2Model::new(vb, &cfg, eos_ids)?;
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -57,7 +57,6 @@ impl<'a> Lfm2GenerateModel<'a> {
|
|||||||
tokenizer,
|
tokenizer,
|
||||||
device,
|
device,
|
||||||
model,
|
model,
|
||||||
eos_token_id,
|
|
||||||
model_name,
|
model_name,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -66,40 +65,28 @@ impl<'a> Lfm2GenerateModel<'a> {
|
|||||||
impl<'a> GenerateModel for Lfm2GenerateModel<'a> {
|
impl<'a> GenerateModel for Lfm2GenerateModel<'a> {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
let mut logits = get_logit_processor(
|
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||||
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
|
let mut ctx = GenerationContext::new(
|
||||||
mes.temperature,
|
mes.temperature,
|
||||||
mes.top_p,
|
mes.top_p,
|
||||||
None,
|
None,
|
||||||
mes.seed.unwrap_or(34562) as u64,
|
seed,
|
||||||
|
input_ids.dim(1)?,
|
||||||
|
sample_len,
|
||||||
|
self.device.clone(),
|
||||||
);
|
);
|
||||||
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
let data = MultiModalData::new(vec![]);
|
||||||
let prompt_tokens = seq_len as u32;
|
generate_generic(
|
||||||
let mut seqlen_offset = 0;
|
&mut self.model,
|
||||||
let mut generate = vec![];
|
&self.tokenizer,
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
input_ids,
|
||||||
for _ in 0..sample_len {
|
data,
|
||||||
let logit = self.model.forward(&input_ids, seqlen_offset)?;
|
&mut ctx,
|
||||||
let logit = logit.squeeze(0)?.squeeze(0)?;
|
|
||||||
let next_token = logits.sample(&logit)?;
|
|
||||||
generate.push(next_token);
|
|
||||||
if next_token == self.eos_token_id {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?;
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
}
|
|
||||||
self.model.clear_cache();
|
|
||||||
let completion_tokens = generate.len() as u32;
|
|
||||||
let decode = self.tokenizer.token_decode(generate)?;
|
|
||||||
let mes = build_completion_response(
|
|
||||||
decode,
|
|
||||||
&self.model_name,
|
&self.model_name,
|
||||||
Some(completion_tokens),
|
)
|
||||||
Some(prompt_tokens),
|
|
||||||
);
|
|
||||||
Ok(mes)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
@@ -115,50 +102,24 @@ impl<'a> GenerateModel for Lfm2GenerateModel<'a> {
|
|||||||
>,
|
>,
|
||||||
> {
|
> {
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
let mut logits = get_logit_processor(
|
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||||
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
|
let data = MultiModalData::new(vec![]);
|
||||||
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
|
let stream = generate_stream_generic(
|
||||||
|
&mut self.model,
|
||||||
|
&self.tokenizer,
|
||||||
|
input_ids,
|
||||||
|
data,
|
||||||
mes.temperature,
|
mes.temperature,
|
||||||
mes.top_p,
|
mes.top_p,
|
||||||
None,
|
None,
|
||||||
mes.seed.unwrap_or(34562) as u64,
|
seed,
|
||||||
);
|
sample_len,
|
||||||
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
false,
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
&self.device,
|
||||||
let mut seqlen_offset = 0;
|
&self.model_name,
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
)?;
|
||||||
let stream = stream! {
|
|
||||||
let mut err_tokens = vec![];
|
|
||||||
for _ in 0..sample_len {
|
|
||||||
let logit = self.model.forward(&input_ids, seqlen_offset)?;
|
|
||||||
let logit = logit.squeeze(0)?.squeeze(0)?;
|
|
||||||
let next_token = logits.sample(&logit)?;
|
|
||||||
let mut decode_ids = vec![];
|
|
||||||
if !err_tokens.is_empty() {
|
|
||||||
decode_ids.extend_from_slice(&err_tokens);
|
|
||||||
}
|
|
||||||
decode_ids.push(next_token);
|
|
||||||
let decode = self.tokenizer.token_decode(decode_ids)?;
|
|
||||||
if decode.contains("�") {
|
|
||||||
err_tokens.push(next_token);
|
|
||||||
if err_tokens.len() > 3 {
|
|
||||||
err_tokens.clear();
|
|
||||||
}
|
|
||||||
input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?;
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
err_tokens.clear();
|
|
||||||
let chunk = build_completion_chunk_response(decode, &self.model_name, None, None);
|
|
||||||
yield Ok(chunk);
|
|
||||||
if next_token == self.eos_token_id {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?;
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
}
|
|
||||||
self.model.clear_cache();
|
|
||||||
};
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,6 +1,9 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::modules::{GateUpDownMLP, QKNormAttention, conv1d_depthwise, get_conv1d},
|
common::{
|
||||||
|
InferenceModel,
|
||||||
|
modules::{GateUpDownMLP, QKNormAttention, conv1d_depthwise, get_conv1d},
|
||||||
|
},
|
||||||
lfm2::config::Lfm2Config,
|
lfm2::config::Lfm2Config,
|
||||||
},
|
},
|
||||||
position_embed::rope::RoPE,
|
position_embed::rope::RoPE,
|
||||||
@@ -279,10 +282,11 @@ impl Lfm2Decoder {
|
|||||||
pub struct Lfm2Model {
|
pub struct Lfm2Model {
|
||||||
model: Lfm2Decoder,
|
model: Lfm2Decoder,
|
||||||
lm_head: Linear,
|
lm_head: Linear,
|
||||||
|
stop_token_ids: Vec<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Lfm2Model {
|
impl Lfm2Model {
|
||||||
pub fn new(vb: VarBuilder, config: &Lfm2Config) -> Result<Self> {
|
pub fn new(vb: VarBuilder, config: &Lfm2Config, eos_ids: Vec<u32>) -> Result<Self> {
|
||||||
let model = Lfm2Decoder::new(vb.pp("model"), config)?;
|
let model = Lfm2Decoder::new(vb.pp("model"), config)?;
|
||||||
let lm_head = if let Some(flag) = config.tie_embedding
|
let lm_head = if let Some(flag) = config.tie_embedding
|
||||||
&& flag
|
&& flag
|
||||||
@@ -300,7 +304,11 @@ impl Lfm2Model {
|
|||||||
Err(_) => Linear::new(model.embed_tokens.embeddings().clone(), None),
|
Err(_) => Linear::new(model.embed_tokens.embeddings().clone(), None),
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
Ok(Self { model, lm_head })
|
Ok(Self {
|
||||||
|
model,
|
||||||
|
lm_head,
|
||||||
|
stop_token_ids: eos_ids,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
@@ -315,3 +323,17 @@ impl Lfm2Model {
|
|||||||
self.model.clear_cache();
|
self.model.clear_cache();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl InferenceModel for Lfm2Model {
|
||||||
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
|
self.forward(input_ids, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clear_cache(&mut self) {
|
||||||
|
self.clear_cache();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn stop_token_ids(&self) -> Vec<u32> {
|
||||||
|
self.stop_token_ids.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+53
-111
@@ -1,9 +1,12 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::common::generate::get_logit_processor,
|
models::common::{
|
||||||
|
MultiModalData,
|
||||||
|
generate::{GenerationContext, generate_generic, generate_stream_generic},
|
||||||
|
},
|
||||||
params::chat::{ChatCompletionParameters, ChatCompletionResponse},
|
params::chat::{ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
@@ -14,12 +17,8 @@ use crate::{
|
|||||||
lfm2vl::{config::Lfm2VLConfig, model::Lfm2VLModel, processor::Lfm2VLProcessor},
|
lfm2vl::{config::Lfm2VLConfig, model::Lfm2VLModel, processor::Lfm2VLProcessor},
|
||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{find_type_files, get_device, get_dtype},
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
|
||||||
get_dtype,
|
|
||||||
},
|
|
||||||
};
|
};
|
||||||
use rocket::async_stream::stream;
|
|
||||||
|
|
||||||
pub struct Lfm2VLGenerateModel<'a> {
|
pub struct Lfm2VLGenerateModel<'a> {
|
||||||
chat_template: ChatTemplate<'a>,
|
chat_template: ChatTemplate<'a>,
|
||||||
@@ -27,7 +26,6 @@ pub struct Lfm2VLGenerateModel<'a> {
|
|||||||
device: Device,
|
device: Device,
|
||||||
model: Lfm2VLModel,
|
model: Lfm2VLModel,
|
||||||
processor: Lfm2VLProcessor,
|
processor: Lfm2VLProcessor,
|
||||||
eos_token_id: u32,
|
|
||||||
model_name: String,
|
model_name: String,
|
||||||
}
|
}
|
||||||
impl<'a> Lfm2VLGenerateModel<'a> {
|
impl<'a> Lfm2VLGenerateModel<'a> {
|
||||||
@@ -43,9 +41,9 @@ impl<'a> Lfm2VLGenerateModel<'a> {
|
|||||||
let model_path = find_type_files(path, "safetensors")?;
|
let model_path = find_type_files(path, "safetensors")?;
|
||||||
let dtype = get_dtype(dtype, &cfg.dtype);
|
let dtype = get_dtype(dtype, &cfg.dtype);
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_path, dtype, &device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_path, dtype, &device)? };
|
||||||
let model = Lfm2VLModel::new(vb, &cfg)?;
|
let eos_ids = vec![gen_cfg.eos_token_id];
|
||||||
|
let model = Lfm2VLModel::new(vb, &cfg, eos_ids)?;
|
||||||
let processor = Lfm2VLProcessor::new(path, dtype, &device)?;
|
let processor = Lfm2VLProcessor::new(path, dtype, &device)?;
|
||||||
let eos_token_id = gen_cfg.eos_token_id;
|
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -57,7 +55,6 @@ impl<'a> Lfm2VLGenerateModel<'a> {
|
|||||||
device,
|
device,
|
||||||
model,
|
model,
|
||||||
processor,
|
processor,
|
||||||
eos_token_id,
|
|
||||||
model_name,
|
model_name,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -66,54 +63,35 @@ impl<'a> Lfm2VLGenerateModel<'a> {
|
|||||||
impl<'a> GenerateModel for Lfm2VLGenerateModel<'a> {
|
impl<'a> GenerateModel for Lfm2VLGenerateModel<'a> {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
let mut logits = get_logit_processor(
|
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 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.temperature,
|
||||||
mes.top_p,
|
mes.top_p,
|
||||||
None,
|
None,
|
||||||
mes.seed.unwrap_or(34562) as u64,
|
seed,
|
||||||
|
input_ids.dim(1)?,
|
||||||
|
sample_len,
|
||||||
|
self.device.clone(),
|
||||||
);
|
);
|
||||||
let (pixel_values, pixel_attention_mask, spatial_shapes, text) =
|
|
||||||
self.processor.process_info(&mes, &mes_render)?;
|
let data_vec = vec![
|
||||||
let mut input_ids = self.tokenizer.text_encode(text, &self.device)?;
|
pixel_values.into(),
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
pixel_attention_mask.into(),
|
||||||
let prompt_tokens = seq_len as u32;
|
spatial_shapes.into(),
|
||||||
let mut seqlen_offset = 0;
|
];
|
||||||
let mut generate = vec![];
|
let data = MultiModalData::new(data_vec);
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
generate_generic(
|
||||||
let mut pixel_values = Some(pixel_values);
|
&mut self.model,
|
||||||
let mut pixel_attention_mask = Some(pixel_attention_mask);
|
&self.tokenizer,
|
||||||
let mut spatial_shapes = Some(spatial_shapes);
|
input_ids,
|
||||||
for _ in 0..sample_len {
|
data,
|
||||||
let logit = self.model.forward(
|
&mut ctx,
|
||||||
&input_ids,
|
|
||||||
pixel_values.as_ref(),
|
|
||||||
pixel_attention_mask.as_ref(),
|
|
||||||
spatial_shapes.as_ref(),
|
|
||||||
seqlen_offset,
|
|
||||||
)?;
|
|
||||||
let logit = logit.squeeze(0)?.squeeze(0)?;
|
|
||||||
let next_token = logits.sample(&logit)?;
|
|
||||||
generate.push(next_token);
|
|
||||||
if next_token == self.eos_token_id {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?;
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
pixel_values = None;
|
|
||||||
pixel_attention_mask = None;
|
|
||||||
spatial_shapes = None;
|
|
||||||
}
|
|
||||||
self.model.clear_cache();
|
|
||||||
let completion_tokens = generate.len() as u32;
|
|
||||||
let decode = self.tokenizer.token_decode(generate)?;
|
|
||||||
let mes = build_completion_response(
|
|
||||||
decode,
|
|
||||||
&self.model_name,
|
&self.model_name,
|
||||||
Some(completion_tokens),
|
)
|
||||||
Some(prompt_tokens),
|
|
||||||
);
|
|
||||||
Ok(mes)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
@@ -129,67 +107,31 @@ impl<'a> GenerateModel for Lfm2VLGenerateModel<'a> {
|
|||||||
>,
|
>,
|
||||||
> {
|
> {
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
let mut logits = get_logit_processor(
|
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.temperature,
|
||||||
mes.top_p,
|
mes.top_p,
|
||||||
None,
|
None,
|
||||||
mes.seed.unwrap_or(34562) as u64,
|
seed,
|
||||||
);
|
sample_len,
|
||||||
let (pixel_values, pixel_attention_mask, spatial_shapes, text) =
|
false,
|
||||||
self.processor.process_info(&mes, &mes_render)?;
|
&self.device,
|
||||||
let mut input_ids = self.tokenizer.text_encode(text, &self.device)?;
|
&self.model_name,
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
|
||||||
let mut pixel_values = Some(pixel_values);
|
|
||||||
let mut pixel_attention_mask = Some(pixel_attention_mask);
|
|
||||||
let mut spatial_shapes = Some(spatial_shapes);
|
|
||||||
let stream = stream! {
|
|
||||||
let mut err_tokens = vec![];
|
|
||||||
for _ in 0..sample_len {
|
|
||||||
let logit = self.model.forward(
|
|
||||||
&input_ids,
|
|
||||||
pixel_values.as_ref(),
|
|
||||||
pixel_attention_mask.as_ref(),
|
|
||||||
spatial_shapes.as_ref(),
|
|
||||||
seqlen_offset,
|
|
||||||
)?;
|
)?;
|
||||||
let logit = logit.squeeze(0)?.squeeze(0)?;
|
|
||||||
let next_token = logits.sample(&logit)?;
|
|
||||||
let mut decode_ids = vec![];
|
|
||||||
if !err_tokens.is_empty() {
|
|
||||||
decode_ids.extend_from_slice(&err_tokens);
|
|
||||||
}
|
|
||||||
decode_ids.push(next_token);
|
|
||||||
let decode = self.tokenizer.token_decode(decode_ids)?;
|
|
||||||
if decode.contains("�") {
|
|
||||||
err_tokens.push(next_token);
|
|
||||||
if err_tokens.len() > 3 {
|
|
||||||
err_tokens.clear();
|
|
||||||
}
|
|
||||||
input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?;
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
pixel_values = None;
|
|
||||||
pixel_attention_mask = None;
|
|
||||||
spatial_shapes = None;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
err_tokens.clear();
|
|
||||||
let chunk = build_completion_chunk_response(decode, &self.model_name, None, None);
|
|
||||||
yield Ok(chunk);
|
|
||||||
if next_token == self.eos_token_id {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?;
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
pixel_values = None;
|
|
||||||
pixel_attention_mask = None;
|
|
||||||
spatial_shapes = None;
|
|
||||||
}
|
|
||||||
self.model.clear_cache();
|
|
||||||
};
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+44
-10
@@ -1,6 +1,9 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::modules::{NaiveAttnTwoLinearMLPBlock, get_layer_norm},
|
common::{
|
||||||
|
InferenceModel,
|
||||||
|
modules::{NaiveAttnTwoLinearMLPBlock, get_layer_norm},
|
||||||
|
},
|
||||||
lfm2::model::Lfm2Decoder,
|
lfm2::model::Lfm2Decoder,
|
||||||
lfm2vl::config::{Lfm2VLConfig, Lfm2VLVisionConfig},
|
lfm2vl::config::{Lfm2VLConfig, Lfm2VLVisionConfig},
|
||||||
},
|
},
|
||||||
@@ -15,11 +18,7 @@ use candle_nn::{Activation, LayerNorm, Linear, Module, VarBuilder, embedding, li
|
|||||||
use num::integer::Roots;
|
use num::integer::Roots;
|
||||||
|
|
||||||
pub struct Siglip2VisionEmbeddings {
|
pub struct Siglip2VisionEmbeddings {
|
||||||
// embed_dim: usize,
|
|
||||||
// patch_size: usize,
|
|
||||||
patch_embedding: Linear,
|
patch_embedding: Linear,
|
||||||
// position_embedding_size: usize,
|
|
||||||
// position_embedding: Embedding,
|
|
||||||
postitional_embeddings: Tensor,
|
postitional_embeddings: Tensor,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -44,11 +43,7 @@ impl Siglip2VisionEmbeddings {
|
|||||||
.permute((2, 0, 1))?
|
.permute((2, 0, 1))?
|
||||||
.unsqueeze(0)?;
|
.unsqueeze(0)?;
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
// embed_dim,
|
|
||||||
// patch_size,
|
|
||||||
patch_embedding,
|
patch_embedding,
|
||||||
// position_embedding_size,
|
|
||||||
// position_embedding,
|
|
||||||
postitional_embeddings,
|
postitional_embeddings,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -254,10 +249,11 @@ pub struct Lfm2VLModel {
|
|||||||
language_model: Lfm2Decoder,
|
language_model: Lfm2Decoder,
|
||||||
lm_head: Linear,
|
lm_head: Linear,
|
||||||
img_id: u32,
|
img_id: u32,
|
||||||
|
stop_token_ids: Vec<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Lfm2VLModel {
|
impl Lfm2VLModel {
|
||||||
pub fn new(vb: VarBuilder, cfg: &Lfm2VLConfig) -> Result<Self> {
|
pub fn new(vb: VarBuilder, cfg: &Lfm2VLConfig, eos_ids: Vec<u32>) -> Result<Self> {
|
||||||
let vb = vb.pp("model");
|
let vb = vb.pp("model");
|
||||||
let vision_tower = Siglip2VisionModel::new(vb.pp("vision_tower"), &cfg.vision_config)?;
|
let vision_tower = Siglip2VisionModel::new(vb.pp("vision_tower"), &cfg.vision_config)?;
|
||||||
let multi_modal_projector =
|
let multi_modal_projector =
|
||||||
@@ -270,6 +266,7 @@ impl Lfm2VLModel {
|
|||||||
language_model,
|
language_model,
|
||||||
lm_head,
|
lm_head,
|
||||||
img_id: cfg.image_token_id,
|
img_id: cfg.image_token_id,
|
||||||
|
stop_token_ids: eos_ids,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -321,3 +318,40 @@ impl Lfm2VLModel {
|
|||||||
self.language_model.clear_cache();
|
self.language_model.clear_cache();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl InferenceModel for Lfm2VLModel {
|
||||||
|
fn forward_initial(
|
||||||
|
&mut self,
|
||||||
|
input_ids: &Tensor,
|
||||||
|
seqlen_offset: usize,
|
||||||
|
data: crate::models::common::MultiModalData,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
if data.data_vec.len() != 3 {
|
||||||
|
return Err(anyhow::anyhow!(
|
||||||
|
"Lfm2VL process data error, must have pixel_values, pixel_attention_mask, spatial_shapes"
|
||||||
|
));
|
||||||
|
}
|
||||||
|
let pixel_values = &data.data_vec[0];
|
||||||
|
let pixel_attention_mask = &data.data_vec[1];
|
||||||
|
let spatial_shapes = &data.data_vec[2];
|
||||||
|
self.forward(
|
||||||
|
input_ids,
|
||||||
|
pixel_values.as_ref(),
|
||||||
|
pixel_attention_mask.as_ref(),
|
||||||
|
spatial_shapes.as_ref(),
|
||||||
|
seqlen_offset,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
|
self.forward(input_ids, None, None, None, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clear_cache(&mut self) {
|
||||||
|
self.clear_cache();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn stop_token_ids(&self) -> Vec<u32> {
|
||||||
|
self.stop_token_ids.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,20 +1,19 @@
|
|||||||
use crate::models::common::generate::get_logit_processor;
|
use crate::models::common::MultiModalData;
|
||||||
|
use crate::models::common::generate::{
|
||||||
|
GenerationContext, generate_generic, generate_stream_generic,
|
||||||
|
};
|
||||||
use crate::params::chat::{
|
use crate::params::chat::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use rocket::async_stream::stream;
|
|
||||||
use rocket::futures::Stream;
|
use rocket::futures::Stream;
|
||||||
|
|
||||||
use crate::models::minicpm4::config::MiniCPM4Config;
|
use crate::models::minicpm4::config::MiniCPM4Config;
|
||||||
use crate::models::minicpm4::model::MiniCPMModel;
|
use crate::models::minicpm4::model::MiniCPMModel;
|
||||||
// use crate::models::GenerateStream;
|
// use crate::models::GenerateStream;
|
||||||
use crate::utils::{
|
use crate::utils::{find_type_files, get_device, get_dtype};
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
|
||||||
get_dtype,
|
|
||||||
};
|
|
||||||
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
||||||
|
|
||||||
pub struct MiniCPMGenerateModel<'a> {
|
pub struct MiniCPMGenerateModel<'a> {
|
||||||
@@ -22,8 +21,6 @@ pub struct MiniCPMGenerateModel<'a> {
|
|||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
minicpm: MiniCPMModel,
|
minicpm: MiniCPMModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
endoftext_id: u32,
|
|
||||||
im_end_id: u32,
|
|
||||||
model_name: String,
|
model_name: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -40,7 +37,8 @@ impl<'a> MiniCPMGenerateModel<'a> {
|
|||||||
let im_end_id = cfg.eos_token_id[1];
|
let im_end_id = cfg.eos_token_id[1];
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
||||||
let minicpm = MiniCPMModel::new(vb, cfg)?;
|
let eos_ids = vec![endoftext_id, im_end_id];
|
||||||
|
let minicpm = MiniCPMModel::new(vb, cfg, eos_ids)?;
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -51,8 +49,6 @@ impl<'a> MiniCPMGenerateModel<'a> {
|
|||||||
tokenizer,
|
tokenizer,
|
||||||
minicpm,
|
minicpm,
|
||||||
device: device.clone(),
|
device: device.clone(),
|
||||||
endoftext_id,
|
|
||||||
im_end_id,
|
|
||||||
model_name,
|
model_name,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -60,33 +56,29 @@ impl<'a> MiniCPMGenerateModel<'a> {
|
|||||||
|
|
||||||
impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
|
||||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let prompt_tokens = seq_len as u32;
|
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
let mut generate = Vec::new();
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(2048);
|
let sample_len = mes.max_tokens.unwrap_or(2048);
|
||||||
for _ in 0..sample_len {
|
let mut ctx = GenerationContext::new(
|
||||||
let logits = self.minicpm.forward_with_cache(&input_ids, seqlen_offset)?;
|
mes.temperature,
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
mes.top_p,
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
None,
|
||||||
generate.push(next_token);
|
seed,
|
||||||
if next_token == self.endoftext_id || next_token == self.im_end_id {
|
input_ids.dim(1)?,
|
||||||
break;
|
sample_len,
|
||||||
}
|
self.device.clone(),
|
||||||
seqlen_offset += seq_len;
|
);
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
let data = MultiModalData::new(vec![]);
|
||||||
}
|
generate_generic(
|
||||||
let num_token = generate.len() as u32;
|
&mut self.minicpm,
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
&self.tokenizer,
|
||||||
self.minicpm.clear_kv_cache();
|
input_ids,
|
||||||
let response =
|
data,
|
||||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
&mut ctx,
|
||||||
Ok(response)
|
&self.model_name,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
&mut self,
|
&mut self,
|
||||||
@@ -100,50 +92,24 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
|||||||
>,
|
>,
|
||||||
> {
|
> {
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
let data = MultiModalData::new(vec![]);
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(512);
|
let sample_len = mes.max_tokens.unwrap_or(512);
|
||||||
let stream = stream! {
|
let stream = generate_stream_generic(
|
||||||
let mut error_tokens = Vec::new();
|
&mut self.minicpm,
|
||||||
for _ in 0..sample_len {
|
&self.tokenizer,
|
||||||
let logits = self.minicpm.forward_with_cache(
|
input_ids,
|
||||||
&input_ids,
|
data,
|
||||||
seqlen_offset,
|
mes.temperature,
|
||||||
|
mes.top_p,
|
||||||
|
None,
|
||||||
|
seed,
|
||||||
|
sample_len,
|
||||||
|
false,
|
||||||
|
&self.device,
|
||||||
|
&self.model_name,
|
||||||
)?;
|
)?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
|
||||||
let mut decode_ids = Vec::new();
|
|
||||||
if !error_tokens.is_empty(){
|
|
||||||
decode_ids.extend_from_slice(&error_tokens);
|
|
||||||
}
|
|
||||||
decode_ids.push(next_token);
|
|
||||||
let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
|
|
||||||
if decoded_token.contains("�") {
|
|
||||||
error_tokens.push(next_token);
|
|
||||||
if error_tokens.len() > 3 {
|
|
||||||
error_tokens.clear();
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
error_tokens.clear();
|
|
||||||
let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
|
||||||
yield Ok(chunk);
|
|
||||||
if next_token == self.endoftext_id || next_token == self.im_end_id {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
|
|
||||||
}
|
|
||||||
self.minicpm.clear_kv_cache();
|
|
||||||
};
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,7 +4,10 @@ use candle_nn::{Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, rms_n
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::modules::{GateUpDownMLP, NaiveAttention},
|
common::{
|
||||||
|
InferenceModel,
|
||||||
|
modules::{GateUpDownMLP, NaiveAttention},
|
||||||
|
},
|
||||||
minicpm4::config::MiniCPM4Config,
|
minicpm4::config::MiniCPM4Config,
|
||||||
},
|
},
|
||||||
position_embed::rope::compute_default_rope_parameters,
|
position_embed::rope::compute_default_rope_parameters,
|
||||||
@@ -207,10 +210,11 @@ pub struct MiniCPMModel {
|
|||||||
norm: RmsNorm,
|
norm: RmsNorm,
|
||||||
rope_emb: MiniCPMLongRoPE,
|
rope_emb: MiniCPMLongRoPE,
|
||||||
lm_head: Linear,
|
lm_head: Linear,
|
||||||
|
stop_token_ids: Vec<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl MiniCPMModel {
|
impl MiniCPMModel {
|
||||||
pub fn new(vb: VarBuilder, cfg: MiniCPM4Config) -> Result<Self> {
|
pub fn new(vb: VarBuilder, cfg: MiniCPM4Config, eos_ids: Vec<u32>) -> Result<Self> {
|
||||||
let vb = vb.pp("model");
|
let vb = vb.pp("model");
|
||||||
let embed_tokens = embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("embed_tokens"))?;
|
let embed_tokens = embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("embed_tokens"))?;
|
||||||
let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
|
let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
|
||||||
@@ -229,10 +233,11 @@ impl MiniCPMModel {
|
|||||||
norm,
|
norm,
|
||||||
rope_emb,
|
rope_emb,
|
||||||
lm_head,
|
lm_head,
|
||||||
|
stop_token_ids: eos_ids,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn forward(&mut self, input_ids: &Tensor, position_id: usize) -> Result<Tensor> {
|
pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
let (bs, seq_len) = input_ids.dims2()?;
|
let (bs, seq_len) = input_ids.dims2()?;
|
||||||
let input_embeds = self
|
let input_embeds = self
|
||||||
.embed_tokens
|
.embed_tokens
|
||||||
@@ -251,7 +256,7 @@ impl MiniCPMModel {
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?;
|
let (cos, sin) = self.rope_emb.forward(seqlen_offset, seq_len)?;
|
||||||
let mut hidden_states = input_embeds;
|
let mut hidden_states = input_embeds;
|
||||||
for decode_layer in &self.layers {
|
for decode_layer in &self.layers {
|
||||||
hidden_states =
|
hidden_states =
|
||||||
@@ -267,7 +272,11 @@ impl MiniCPMModel {
|
|||||||
Ok(logits)
|
Ok(logits)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn forward_with_cache(&mut self, input_ids: &Tensor, position_id: usize) -> Result<Tensor> {
|
pub fn forward_with_cache(
|
||||||
|
&mut self,
|
||||||
|
input_ids: &Tensor,
|
||||||
|
seqlen_offset: usize,
|
||||||
|
) -> Result<Tensor> {
|
||||||
let (bs, seq_len) = input_ids.dims2()?;
|
let (bs, seq_len) = input_ids.dims2()?;
|
||||||
let input_embeds = self
|
let input_embeds = self
|
||||||
.embed_tokens
|
.embed_tokens
|
||||||
@@ -285,7 +294,7 @@ impl MiniCPMModel {
|
|||||||
)?)
|
)?)
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?;
|
let (cos, sin) = self.rope_emb.forward(seqlen_offset, seq_len)?;
|
||||||
let mut hidden_states = input_embeds;
|
let mut hidden_states = input_embeds;
|
||||||
for decode_layer in &mut self.layers {
|
for decode_layer in &mut self.layers {
|
||||||
hidden_states = decode_layer.forward_with_cache(
|
hidden_states = decode_layer.forward_with_cache(
|
||||||
@@ -311,3 +320,17 @@ impl MiniCPMModel {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl InferenceModel for MiniCPMModel {
|
||||||
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
|
self.forward_with_cache(input_ids, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clear_cache(&mut self) {
|
||||||
|
self.clear_kv_cache();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn stop_token_ids(&self) -> Vec<u32> {
|
||||||
|
self.stop_token_ids.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,21 +1,20 @@
|
|||||||
use crate::models::common::generate::get_logit_processor;
|
use crate::models::common::MultiModalData;
|
||||||
|
use crate::models::common::generate::{
|
||||||
|
GenerationContext, generate_generic, generate_stream_generic,
|
||||||
|
};
|
||||||
use crate::params::chat::{
|
use crate::params::chat::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
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::async_stream::stream;
|
|
||||||
use rocket::futures::Stream;
|
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::{
|
use crate::utils::{find_type_files, get_device, get_dtype};
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
|
||||||
get_dtype,
|
|
||||||
};
|
|
||||||
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
||||||
|
|
||||||
pub struct PaddleOCRVLGenerateModel<'a> {
|
pub struct PaddleOCRVLGenerateModel<'a> {
|
||||||
@@ -25,7 +24,6 @@ pub struct PaddleOCRVLGenerateModel<'a> {
|
|||||||
paddleocr_vl: PaddleOCRVLModel,
|
paddleocr_vl: PaddleOCRVLModel,
|
||||||
cfg: PaddleOCRVLConfig,
|
cfg: PaddleOCRVLConfig,
|
||||||
device: Device,
|
device: Device,
|
||||||
end_token_id: u32,
|
|
||||||
model_name: String,
|
model_name: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -42,10 +40,9 @@ impl<'a> PaddleOCRVLGenerateModel<'a> {
|
|||||||
let processor_cfg: PaddleOCRVLPreprocessorConfig =
|
let processor_cfg: PaddleOCRVLPreprocessorConfig =
|
||||||
serde_json::from_slice(&std::fs::read(processor_cfg_path)?)?;
|
serde_json::from_slice(&std::fs::read(processor_cfg_path)?)?;
|
||||||
let pre_processor = PaddleOCRVLProcessor::new(processor_cfg, device, dtype)?;
|
let pre_processor = PaddleOCRVLProcessor::new(processor_cfg, device, dtype)?;
|
||||||
let end_token_id = 2;
|
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
||||||
let paddleocr_vl = PaddleOCRVLModel::new(cfg.clone(), vb)?;
|
let paddleocr_vl = PaddleOCRVLModel::new(cfg.clone(), vb, vec![2])?;
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -58,7 +55,6 @@ impl<'a> PaddleOCRVLGenerateModel<'a> {
|
|||||||
paddleocr_vl,
|
paddleocr_vl,
|
||||||
cfg,
|
cfg,
|
||||||
device: device.clone(),
|
device: device.clone(),
|
||||||
end_token_id,
|
|
||||||
model_name,
|
model_name,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -66,53 +62,43 @@ impl<'a> PaddleOCRVLGenerateModel<'a> {
|
|||||||
|
|
||||||
impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
|
impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
|
||||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
let (replace_text, mut pixel_values, mut 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 mut input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
|
let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
|
||||||
let prompt_tokens = seq_len as u32;
|
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
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)?;
|
||||||
|
|
||||||
let mut cache_position = Tensor::ones_like(&input_ids.i(0)?)?
|
let cache_position = Tensor::ones_like(&input_ids.i(0)?)?
|
||||||
.to_dtype(candle_core::DType::F64)?
|
.to_dtype(candle_core::DType::F64)?
|
||||||
.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 mut generate = Vec::new();
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
for _ in 0..sample_len {
|
let mut ctx = GenerationContext::new(
|
||||||
let logits = self.paddleocr_vl.forward(
|
mes.temperature,
|
||||||
&input_ids,
|
mes.top_p,
|
||||||
pixel_values.as_ref(),
|
None,
|
||||||
image_grid_thw.as_ref(),
|
seed,
|
||||||
&image_mask,
|
input_ids.dim(1)?,
|
||||||
Some(&cache_position),
|
sample_len,
|
||||||
seqlen_offset,
|
self.device.clone(),
|
||||||
)?;
|
);
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
let data_vec = vec![
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
pixel_values,
|
||||||
generate.push(next_token);
|
image_grid_thw,
|
||||||
if next_token == self.end_token_id {
|
image_mask.into(),
|
||||||
break;
|
cache_position.into(),
|
||||||
}
|
];
|
||||||
seqlen_offset += seq_len;
|
let data = MultiModalData::new(data_vec);
|
||||||
seq_len = 1;
|
generate_generic(
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
&mut self.paddleocr_vl,
|
||||||
cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
&self.tokenizer,
|
||||||
pixel_values = None;
|
input_ids,
|
||||||
image_grid_thw = None;
|
data,
|
||||||
}
|
&mut ctx,
|
||||||
let num_token = generate.len() as u32;
|
&self.model_name,
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
)
|
||||||
self.paddleocr_vl.clear_kv_cache();
|
|
||||||
let response =
|
|
||||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
|
||||||
Ok(response)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
@@ -126,72 +112,41 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
|
|||||||
+ '_,
|
+ '_,
|
||||||
>,
|
>,
|
||||||
> {
|
> {
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
|
||||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
|
||||||
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 mut input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
|
let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
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)?;
|
||||||
|
|
||||||
let mut cache_position = Tensor::ones_like(&input_ids.i(0)?)?
|
let cache_position = Tensor::ones_like(&input_ids.i(0)?)?
|
||||||
.to_dtype(candle_core::DType::F64)?
|
.to_dtype(candle_core::DType::F64)?
|
||||||
.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 sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
let stream = stream! {
|
let data_vec = vec![
|
||||||
let mut error_tokens = Vec::new();
|
|
||||||
let mut pixel_values = pixel_values.as_ref();
|
|
||||||
let mut image_grid_thw = image_grid_thw.as_ref();
|
|
||||||
for _ in 0..sample_len {
|
|
||||||
let logits = self.paddleocr_vl.forward(
|
|
||||||
&input_ids,
|
|
||||||
pixel_values,
|
pixel_values,
|
||||||
image_grid_thw,
|
image_grid_thw,
|
||||||
&image_mask,
|
image_mask.into(),
|
||||||
Some(&cache_position),
|
cache_position.into(),
|
||||||
seqlen_offset,
|
];
|
||||||
|
let data = MultiModalData::new(data_vec);
|
||||||
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
|
let stream = generate_stream_generic(
|
||||||
|
&mut self.paddleocr_vl,
|
||||||
|
&self.tokenizer,
|
||||||
|
input_ids,
|
||||||
|
data,
|
||||||
|
mes.temperature,
|
||||||
|
mes.top_p,
|
||||||
|
None,
|
||||||
|
seed,
|
||||||
|
sample_len,
|
||||||
|
false,
|
||||||
|
&self.device,
|
||||||
|
&self.model_name,
|
||||||
)?;
|
)?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
|
||||||
let mut decode_ids = Vec::new();
|
|
||||||
if !error_tokens.is_empty() {
|
|
||||||
decode_ids.extend_from_slice(&error_tokens);
|
|
||||||
}
|
|
||||||
decode_ids.push(next_token);
|
|
||||||
let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
|
|
||||||
if decoded_token.contains("�") {
|
|
||||||
error_tokens.push(next_token);
|
|
||||||
if error_tokens.len() > 3 {
|
|
||||||
error_tokens.clear();
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
|
||||||
pixel_values = None;
|
|
||||||
image_grid_thw = None;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
error_tokens.clear();
|
|
||||||
let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
|
||||||
yield Ok(chunk);
|
|
||||||
if next_token == self.end_token_id {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
|
||||||
pixel_values = None;
|
|
||||||
image_grid_thw = None;
|
|
||||||
}
|
|
||||||
self.paddleocr_vl.clear_kv_cache();
|
|
||||||
};
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,18 +8,23 @@ use num::integer::Roots;
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::modules::{
|
common::{
|
||||||
|
InferenceModel,
|
||||||
|
modules::{
|
||||||
NaiveAttnGateUpDownMLPBlock, NaiveAttnTwoLinearMLPBlock, get_conv2d, get_layer_norm,
|
NaiveAttnGateUpDownMLPBlock, NaiveAttnTwoLinearMLPBlock, get_conv2d, get_layer_norm,
|
||||||
},
|
},
|
||||||
|
},
|
||||||
paddleocr_vl::config::{
|
paddleocr_vl::config::{
|
||||||
PaddleOCRVLConfig, PaddleOCRVLRopeScalingConfig, PaddleOCRVLVisionConfig,
|
PaddleOCRVLConfig, PaddleOCRVLRopeScalingConfig, PaddleOCRVLVisionConfig,
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
position_embed::rope::{Qwen2_5VLTextRotaryEmbedding, Qwen2_5VisionRotaryEmbedding},
|
position_embed::rope::{Qwen2_5VLTextRotaryEmbedding, Qwen2_5VisionRotaryEmbedding},
|
||||||
utils::interpolate::interpolate_bilinear,
|
utils::{
|
||||||
utils::tensor_utils::{
|
interpolate::interpolate_bilinear,
|
||||||
get_vision_next_indices, masked_scatter_dim0, nonzero_index, prepare_causal_attention_mask,
|
tensor_utils::{
|
||||||
zero_index,
|
get_vision_next_indices, masked_scatter_dim0, nonzero_index,
|
||||||
|
prepare_causal_attention_mask, zero_index,
|
||||||
|
},
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -412,10 +417,11 @@ pub struct PaddleOCRVLModel {
|
|||||||
pub cfg: PaddleOCRVLConfig,
|
pub cfg: PaddleOCRVLConfig,
|
||||||
lm_head: Linear,
|
lm_head: Linear,
|
||||||
rope_deltas: Option<Tensor>,
|
rope_deltas: Option<Tensor>,
|
||||||
|
stop_token_ids: Vec<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl PaddleOCRVLModel {
|
impl PaddleOCRVLModel {
|
||||||
pub fn new(cfg: PaddleOCRVLConfig, vb: VarBuilder) -> Result<Self> {
|
pub fn new(cfg: PaddleOCRVLConfig, vb: VarBuilder, eos_ids: Vec<u32>) -> Result<Self> {
|
||||||
let mlp_ar = Projector::new(vb.pp("mlp_AR"), &cfg)?;
|
let mlp_ar = Projector::new(vb.pp("mlp_AR"), &cfg)?;
|
||||||
let visual = SiglipVisionModel::new(vb.pp("visual"), &cfg.vision_config)?;
|
let visual = SiglipVisionModel::new(vb.pp("visual"), &cfg.vision_config)?;
|
||||||
let model = Ernie4_5Model::new(vb.pp("model"), &cfg)?;
|
let model = Ernie4_5Model::new(vb.pp("model"), &cfg)?;
|
||||||
@@ -433,6 +439,7 @@ impl PaddleOCRVLModel {
|
|||||||
cfg,
|
cfg,
|
||||||
lm_head,
|
lm_head,
|
||||||
rope_deltas: None,
|
rope_deltas: None,
|
||||||
|
stop_token_ids: eos_ids,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -662,13 +669,14 @@ impl PaddleOCRVLModel {
|
|||||||
input_ids: &Tensor,
|
input_ids: &Tensor,
|
||||||
pixel_values: Option<&Tensor>,
|
pixel_values: Option<&Tensor>,
|
||||||
image_grid_thw: Option<&Tensor>,
|
image_grid_thw: Option<&Tensor>,
|
||||||
image_mask: &Tensor,
|
image_mask: Option<&Tensor>,
|
||||||
cache_position: Option<&Tensor>,
|
cache_position: Option<&Tensor>,
|
||||||
seqlen_offset: usize,
|
seqlen_offset: usize,
|
||||||
) -> Result<Tensor> {
|
) -> Result<Tensor> {
|
||||||
let mut inputs_embeds = self.model.embed_tokens.forward(input_ids)?;
|
let mut inputs_embeds = self.model.embed_tokens.forward(input_ids)?;
|
||||||
if let Some(pixel_values) = pixel_values
|
if let Some(pixel_values) = pixel_values
|
||||||
&& let Some(image_grid_thw) = image_grid_thw
|
&& let Some(image_grid_thw) = image_grid_thw
|
||||||
|
&& let Some(image_mask) = image_mask
|
||||||
{
|
{
|
||||||
let pixel_values = pixel_values.unsqueeze(0)?;
|
let pixel_values = pixel_values.unsqueeze(0)?;
|
||||||
let mut siglip_position_ids = vec![];
|
let mut siglip_position_ids = vec![];
|
||||||
@@ -716,6 +724,15 @@ impl PaddleOCRVLModel {
|
|||||||
.broadcast_add(rope_deltas)?
|
.broadcast_add(rope_deltas)?
|
||||||
.contiguous()?
|
.contiguous()?
|
||||||
.to_dtype(candle_core::DType::U32)?
|
.to_dtype(candle_core::DType::U32)?
|
||||||
|
} else if let Some(rope_deltas) = &self.rope_deltas {
|
||||||
|
let cache_position =
|
||||||
|
Tensor::from_vec(vec![seqlen_offset as u32], 1, inputs_embeds.device())?;
|
||||||
|
cache_position
|
||||||
|
.i(0)?
|
||||||
|
.to_dtype(rope_deltas.dtype())?
|
||||||
|
.broadcast_add(rope_deltas)?
|
||||||
|
.contiguous()?
|
||||||
|
.to_dtype(candle_core::DType::U32)?
|
||||||
} else {
|
} else {
|
||||||
Tensor::zeros(1, inputs_embeds.dtype(), inputs_embeds.device())?
|
Tensor::zeros(1, inputs_embeds.dtype(), inputs_embeds.device())?
|
||||||
};
|
};
|
||||||
@@ -740,3 +757,42 @@ impl PaddleOCRVLModel {
|
|||||||
self.model.clear_kv_cache();
|
self.model.clear_kv_cache();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl InferenceModel for PaddleOCRVLModel {
|
||||||
|
fn forward_initial(
|
||||||
|
&mut self,
|
||||||
|
input_ids: &Tensor,
|
||||||
|
seqlen_offset: usize,
|
||||||
|
data: crate::models::common::MultiModalData,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
if data.data_vec.len() != 4 {
|
||||||
|
return Err(anyhow::anyhow!(
|
||||||
|
"Lfm2VL process data error, must have pixel_values, image_grid_thw, image_mask, cache_position"
|
||||||
|
));
|
||||||
|
}
|
||||||
|
let pixel_values = &data.data_vec[0];
|
||||||
|
let image_grid_thw = &data.data_vec[1];
|
||||||
|
let image_mask = &data.data_vec[2];
|
||||||
|
let cache_position = &data.data_vec[3];
|
||||||
|
self.forward(
|
||||||
|
input_ids,
|
||||||
|
pixel_values.as_ref(),
|
||||||
|
image_grid_thw.as_ref(),
|
||||||
|
image_mask.as_ref(),
|
||||||
|
cache_position.as_ref(),
|
||||||
|
seqlen_offset,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
|
self.forward(input_ids, None, None, None, None, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clear_cache(&mut self) {
|
||||||
|
self.clear_kv_cache();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn stop_token_ids(&self) -> Vec<u32> {
|
||||||
|
self.stop_token_ids.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,7 +1,12 @@
|
|||||||
|
use std::time::Instant;
|
||||||
|
|
||||||
use crate::models::common::generate::get_logit_processor;
|
use crate::models::common::generate::get_logit_processor;
|
||||||
use crate::params::chat::{
|
use crate::params::chat::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||||
};
|
};
|
||||||
|
use crate::utils::response_utils::{
|
||||||
|
build_chunk_response_with_usage, build_completion_response_with_time,
|
||||||
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{D, DType, Device, IndexOp, Tensor};
|
use candle_core::{D, DType, Device, IndexOp, Tensor};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
@@ -10,8 +15,7 @@ use rocket::futures::Stream;
|
|||||||
|
|
||||||
use crate::models::qwen2_5vl::config::Qwen2_5VLConfig;
|
use crate::models::qwen2_5vl::config::Qwen2_5VLConfig;
|
||||||
use crate::utils::{
|
use crate::utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
find_type_files, get_device, get_dtype, response_utils::build_completion_chunk_response,
|
||||||
get_dtype,
|
|
||||||
};
|
};
|
||||||
use crate::{
|
use crate::{
|
||||||
chat_template::ChatTemplate,
|
chat_template::ChatTemplate,
|
||||||
@@ -94,7 +98,10 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
|||||||
|
|
||||||
let mut generate = Vec::new();
|
let mut generate = Vec::new();
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
|
let mut prompt_secs = 0.0f64;
|
||||||
|
let mut completion_secs = 0.0f64;
|
||||||
for _ in 0..sample_len {
|
for _ in 0..sample_len {
|
||||||
|
let i_start = Instant::now();
|
||||||
let logits = self.qwen2_5_vl.forward(
|
let logits = self.qwen2_5_vl.forward(
|
||||||
&input_ids,
|
&input_ids,
|
||||||
pixel_values,
|
pixel_values,
|
||||||
@@ -108,6 +115,12 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
|||||||
)?;
|
)?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
let next_token = logit_processor.sample(&logits)?;
|
||||||
|
let i_duration = i_start.elapsed();
|
||||||
|
if seqlen_offset == 0 {
|
||||||
|
prompt_secs += i_duration.as_secs_f64();
|
||||||
|
} else {
|
||||||
|
completion_secs += i_duration.as_secs_f64();
|
||||||
|
};
|
||||||
generate.push(next_token);
|
generate.push(next_token);
|
||||||
if next_token == self.endoftext_id || next_token == self.im_end_id {
|
if next_token == self.endoftext_id || next_token == self.im_end_id {
|
||||||
break;
|
break;
|
||||||
@@ -124,8 +137,14 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
|||||||
let num_token = generate.len() as u32;
|
let num_token = generate.len() as u32;
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
let res = self.tokenizer.token_decode(generate)?;
|
||||||
self.qwen2_5_vl.clear_kv_cache();
|
self.qwen2_5_vl.clear_kv_cache();
|
||||||
let response =
|
let response = build_completion_response_with_time(
|
||||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
res,
|
||||||
|
&self.model_name,
|
||||||
|
num_token.into(),
|
||||||
|
completion_secs.into(),
|
||||||
|
prompt_tokens.into(),
|
||||||
|
prompt_secs.into(),
|
||||||
|
);
|
||||||
Ok(response)
|
Ok(response)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -148,6 +167,10 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
|||||||
.tokenizer
|
.tokenizer
|
||||||
.text_encode(input.replace_text.clone(), &self.device)?;
|
.text_encode(input.replace_text.clone(), &self.device)?;
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
let mut seq_len = input_ids.dim(1)?;
|
||||||
|
let prompt_tokens = seq_len as u32;
|
||||||
|
let mut prompt_secs = 0.0f64;
|
||||||
|
let mut completion_tokens = 0u32;
|
||||||
|
let mut completion_secs = 0.0f64;
|
||||||
let mut seqlen_offset = 0;
|
let mut seqlen_offset = 0;
|
||||||
let mut mask = Tensor::ones_like(&input_ids)?;
|
let mut mask = Tensor::ones_like(&input_ids)?;
|
||||||
let mut cache_position = Tensor::ones_like(&input_ids.i(0)?)?
|
let mut cache_position = Tensor::ones_like(&input_ids.i(0)?)?
|
||||||
@@ -166,6 +189,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
|||||||
let mut tool_call_id = None;
|
let mut tool_call_id = None;
|
||||||
let mut tool_call_content = String::new();
|
let mut tool_call_content = String::new();
|
||||||
for _ in 0..sample_len {
|
for _ in 0..sample_len {
|
||||||
|
let i_start = Instant::now();
|
||||||
let logits = self.qwen2_5_vl.forward(
|
let logits = self.qwen2_5_vl.forward(
|
||||||
&input_ids,
|
&input_ids,
|
||||||
pixel_values,
|
pixel_values,
|
||||||
@@ -179,6 +203,13 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
|||||||
)?;
|
)?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
let next_token = logit_processor.sample(&logits)?;
|
||||||
|
completion_tokens += 1;
|
||||||
|
let i_duration = i_start.elapsed();
|
||||||
|
if seqlen_offset == 0 {
|
||||||
|
prompt_secs += i_duration.as_secs_f64();
|
||||||
|
} else {
|
||||||
|
completion_secs += i_duration.as_secs_f64();
|
||||||
|
};
|
||||||
let mut decode_ids = Vec::new();
|
let mut decode_ids = Vec::new();
|
||||||
if !error_tokens.is_empty() {
|
if !error_tokens.is_empty() {
|
||||||
decode_ids.extend_from_slice(&error_tokens);
|
decode_ids.extend_from_slice(&error_tokens);
|
||||||
@@ -249,9 +280,8 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
|
||||||
// yield Ok(chunk);
|
|
||||||
if next_token == self.endoftext_id || next_token == self.im_end_id {
|
if next_token == self.endoftext_id || next_token == self.im_end_id {
|
||||||
|
yield Ok(build_chunk_response_with_usage(&self.model_name, completion_tokens.into(), completion_secs.into(), prompt_tokens.into(), prompt_secs.into()));
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
seqlen_offset += seq_len;
|
seqlen_offset += seq_len;
|
||||||
|
|||||||
@@ -1,20 +1,18 @@
|
|||||||
use crate::models::common::generate::get_logit_processor;
|
use crate::models::common::MultiModalData;
|
||||||
|
use crate::models::common::generate::{
|
||||||
|
GenerationContext, generate_generic, generate_stream_generic,
|
||||||
|
};
|
||||||
use crate::params::chat::{
|
use crate::params::chat::{
|
||||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use rocket::async_stream::stream;
|
|
||||||
use rocket::futures::Stream;
|
use rocket::futures::Stream;
|
||||||
|
|
||||||
use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig};
|
use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig};
|
||||||
use crate::models::qwen3::model::Qwen3Model;
|
use crate::models::qwen3::model::Qwen3Model;
|
||||||
// use crate::models::GenerateStream;
|
use crate::utils::{find_type_files, get_device, get_dtype};
|
||||||
use crate::utils::{
|
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
|
||||||
get_dtype,
|
|
||||||
};
|
|
||||||
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
||||||
|
|
||||||
pub struct Qwen3GenerateModel<'a> {
|
pub struct Qwen3GenerateModel<'a> {
|
||||||
@@ -22,8 +20,6 @@ pub struct Qwen3GenerateModel<'a> {
|
|||||||
tokenizer: TokenizerModel,
|
tokenizer: TokenizerModel,
|
||||||
qwen3: Qwen3Model,
|
qwen3: Qwen3Model,
|
||||||
device: Device,
|
device: Device,
|
||||||
eos_token_id1: u32,
|
|
||||||
eos_token_id2: u32,
|
|
||||||
generation_config: Qwen3GenerationConfig,
|
generation_config: Qwen3GenerationConfig,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
}
|
}
|
||||||
@@ -39,10 +35,11 @@ impl<'a> Qwen3GenerateModel<'a> {
|
|||||||
let dtype = get_dtype(dtype, cfg_dtype);
|
let dtype = get_dtype(dtype, cfg_dtype);
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
||||||
let qwen3 = Qwen3Model::new(&cfg, vb)?;
|
|
||||||
let generation_config_path = path.to_string() + "/generation_config.json";
|
let generation_config_path = path.to_string() + "/generation_config.json";
|
||||||
let generation_config: Qwen3GenerationConfig =
|
let generation_config: Qwen3GenerationConfig =
|
||||||
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
||||||
|
let qwen3 = Qwen3Model::new(&cfg, vb, generation_config.eos_token_id.clone())?;
|
||||||
|
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -53,8 +50,6 @@ impl<'a> Qwen3GenerateModel<'a> {
|
|||||||
tokenizer,
|
tokenizer,
|
||||||
qwen3,
|
qwen3,
|
||||||
device: device.clone(),
|
device: device.clone(),
|
||||||
eos_token_id1: generation_config.eos_token_id[0] as u32,
|
|
||||||
eos_token_id2: generation_config.eos_token_id[1] as u32,
|
|
||||||
generation_config,
|
generation_config,
|
||||||
model_name,
|
model_name,
|
||||||
})
|
})
|
||||||
@@ -69,38 +64,28 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
|
|||||||
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
||||||
let top_k = self.generation_config.top_k;
|
let top_k = self.generation_config.top_k;
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let mut logit_processor =
|
|
||||||
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
|
||||||
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
// let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
|
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||||
// let mes_render = self
|
|
||||||
// .chat_template
|
|
||||||
// .apply_chat_temp_think(&mes, enable_thinking)?;
|
|
||||||
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
|
||||||
let prompt_tokens = seq_len as u32;
|
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
let mut generate = Vec::new();
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(2048);
|
let sample_len = mes.max_tokens.unwrap_or(2048);
|
||||||
for _ in 0..sample_len {
|
let mut ctx = GenerationContext::new(
|
||||||
let logits = self.qwen3.forward(Some(&input_ids), None, seqlen_offset)?;
|
temperature.into(),
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
top_p.into(),
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
top_k.into(),
|
||||||
generate.push(next_token);
|
seed,
|
||||||
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
input_ids.dim(1)?,
|
||||||
break;
|
sample_len,
|
||||||
}
|
self.device.clone(),
|
||||||
seqlen_offset += seq_len;
|
);
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
let data = MultiModalData::new(vec![]);
|
||||||
}
|
generate_generic(
|
||||||
let num_token = generate.len() as u32;
|
&mut self.qwen3,
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
&self.tokenizer,
|
||||||
self.qwen3.clear_kv_cache();
|
input_ids,
|
||||||
let response =
|
data,
|
||||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
&mut ctx,
|
||||||
Ok(response)
|
&self.model_name,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
&mut self,
|
&mut self,
|
||||||
@@ -119,56 +104,25 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
|
|||||||
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
||||||
let top_k = self.generation_config.top_k;
|
let top_k = self.generation_config.top_k;
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let mut logit_processor =
|
|
||||||
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
// let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
|
let in_reasoning = mes_render.ends_with("<think>\n");
|
||||||
// let mes_render = self
|
let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||||
// .chat_template
|
let data = MultiModalData::new(vec![]);
|
||||||
// .apply_chat_temp_think(&mes, enable_thinking)?;
|
|
||||||
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(512);
|
let sample_len = mes.max_tokens.unwrap_or(512);
|
||||||
let stream = stream! {
|
let stream = generate_stream_generic(
|
||||||
let mut error_tokens = Vec::new();
|
&mut self.qwen3,
|
||||||
for _ in 0..sample_len {
|
&self.tokenizer,
|
||||||
let logits = self.qwen3.forward(
|
input_ids,
|
||||||
Some(&input_ids),
|
data,
|
||||||
None,
|
temperature.into(),
|
||||||
seqlen_offset,
|
top_p.into(),
|
||||||
|
top_k.into(),
|
||||||
|
seed,
|
||||||
|
sample_len,
|
||||||
|
in_reasoning,
|
||||||
|
&self.device,
|
||||||
|
&self.model_name,
|
||||||
)?;
|
)?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
|
||||||
let mut decode_ids = Vec::new();
|
|
||||||
if !error_tokens.is_empty(){
|
|
||||||
decode_ids.extend_from_slice(&error_tokens);
|
|
||||||
}
|
|
||||||
decode_ids.push(next_token);
|
|
||||||
let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
|
|
||||||
if decoded_token.contains("�") {
|
|
||||||
error_tokens.push(next_token);
|
|
||||||
if error_tokens.len() > 3 {
|
|
||||||
error_tokens.clear();
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
error_tokens.clear();
|
|
||||||
let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
|
||||||
yield Ok(chunk);
|
|
||||||
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
|
|
||||||
}
|
|
||||||
self.qwen3.clear_kv_cache();
|
|
||||||
};
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,7 +6,10 @@ use candle_nn::{
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::modules::{GateUpDownMLP, QKNormAttention},
|
common::{
|
||||||
|
InferenceModel,
|
||||||
|
modules::{GateUpDownMLP, QKNormAttention},
|
||||||
|
},
|
||||||
qwen3::config::Qwen3Config,
|
qwen3::config::Qwen3Config,
|
||||||
},
|
},
|
||||||
position_embed::rope::RoPE,
|
position_embed::rope::RoPE,
|
||||||
@@ -94,10 +97,11 @@ pub struct Qwen3Model {
|
|||||||
norm: RmsNorm,
|
norm: RmsNorm,
|
||||||
rotary_emb: RoPE,
|
rotary_emb: RoPE,
|
||||||
lm_head: Linear,
|
lm_head: Linear,
|
||||||
|
stop_token_ids: Vec<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Qwen3Model {
|
impl Qwen3Model {
|
||||||
pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result<Self> {
|
pub fn new(config: &Qwen3Config, vb: VarBuilder, eos_ids: Vec<u32>) -> Result<Self> {
|
||||||
let vb = vb.pp("model");
|
let vb = vb.pp("model");
|
||||||
let vocab_size = config.vocab_size;
|
let vocab_size = config.vocab_size;
|
||||||
let embed_tokens = embedding(vocab_size, config.hidden_size, vb.pp("embed_tokens"))?;
|
let embed_tokens = embedding(vocab_size, config.hidden_size, vb.pp("embed_tokens"))?;
|
||||||
@@ -121,6 +125,7 @@ impl Qwen3Model {
|
|||||||
norm,
|
norm,
|
||||||
rotary_emb,
|
rotary_emb,
|
||||||
lm_head,
|
lm_head,
|
||||||
|
stop_token_ids: eos_ids,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
pub fn forward(
|
pub fn forward(
|
||||||
@@ -178,3 +183,17 @@ impl Qwen3Model {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl InferenceModel for Qwen3Model {
|
||||||
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
|
self.forward(input_ids.into(), None, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clear_cache(&mut self) {
|
||||||
|
self.clear_kv_cache();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn stop_token_ids(&self) -> Vec<u32> {
|
||||||
|
self.stop_token_ids.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+56
-176
@@ -1,11 +1,13 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::common::generate::get_logit_processor,
|
models::common::{
|
||||||
|
MultiModalData,
|
||||||
|
generate::{GenerationContext, generate_generic, generate_stream_generic},
|
||||||
|
},
|
||||||
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{DType, Device, Tensor, quantized::gguf_file};
|
use candle_core::{DType, Device, quantized::gguf_file};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use rocket::async_stream::stream;
|
|
||||||
use rocket::futures::Stream;
|
use rocket::futures::Stream;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
@@ -17,10 +19,7 @@ use crate::{
|
|||||||
qwen3vl::processor::Qwen3VLProcessor,
|
qwen3vl::processor::Qwen3VLProcessor,
|
||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{find_type_files, get_device, get_dtype},
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
|
||||||
get_dtype,
|
|
||||||
},
|
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct Qwen3_5GenerateModel<'a> {
|
pub struct Qwen3_5GenerateModel<'a> {
|
||||||
@@ -29,10 +28,9 @@ pub struct Qwen3_5GenerateModel<'a> {
|
|||||||
pre_processor: Option<Qwen3VLProcessor>,
|
pre_processor: Option<Qwen3VLProcessor>,
|
||||||
qwen3_5: Qwen3_5Model,
|
qwen3_5: Qwen3_5Model,
|
||||||
device: Device,
|
device: Device,
|
||||||
eos_token_id: u32,
|
|
||||||
model_name: String,
|
model_name: String,
|
||||||
repeat_penalty: f32,
|
// repeat_penalty: f32, // TODO
|
||||||
repeat_last_n: usize,
|
// repeat_last_n: usize,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<'a> Qwen3_5GenerateModel<'a> {
|
impl<'a> Qwen3_5GenerateModel<'a> {
|
||||||
@@ -51,8 +49,8 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
|||||||
let pre_processor = Qwen3VLProcessor::new(path, &device, dtype)?;
|
let pre_processor = Qwen3VLProcessor::new(path, &device, dtype)?;
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
||||||
let eos_token_id = cfg.text_config.eos_token_id;
|
let eos_ids = vec![cfg.text_config.eos_token_id];
|
||||||
let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg)?;
|
let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg, eos_ids)?;
|
||||||
|
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
chat_template,
|
chat_template,
|
||||||
@@ -60,10 +58,9 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
|||||||
pre_processor: Some(pre_processor),
|
pre_processor: Some(pre_processor),
|
||||||
qwen3_5,
|
qwen3_5,
|
||||||
device,
|
device,
|
||||||
eos_token_id,
|
|
||||||
model_name: model_name.to_string(),
|
model_name: model_name.to_string(),
|
||||||
repeat_penalty: 1.01,
|
// repeat_penalty: 1.01,
|
||||||
repeat_last_n: 64,
|
// repeat_last_n: 64,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -105,7 +102,9 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
|||||||
let eos_token_id = model_gguf
|
let eos_token_id = model_gguf
|
||||||
.get_matedata("tokenizer.ggml.eos_token_id")?
|
.get_matedata("tokenizer.ggml.eos_token_id")?
|
||||||
.to_u32()?;
|
.to_u32()?;
|
||||||
let qwen3_5 = Qwen3_5Model::new_from_gguf(&mut model_gguf, mmproj_gguf.as_mut(), &device)?;
|
let eos_ids = vec![eos_token_id];
|
||||||
|
let qwen3_5 =
|
||||||
|
Qwen3_5Model::new_from_gguf(&mut model_gguf, mmproj_gguf.as_mut(), &device, eos_ids)?;
|
||||||
let stem = std::path::Path::new(model_file)
|
let stem = std::path::Path::new(model_file)
|
||||||
.file_stem() // 获取文件名主干(不含扩展名)
|
.file_stem() // 获取文件名主干(不含扩展名)
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -116,11 +115,9 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
|||||||
pre_processor,
|
pre_processor,
|
||||||
qwen3_5,
|
qwen3_5,
|
||||||
device,
|
device,
|
||||||
// eos_token_id: 248044,
|
|
||||||
eos_token_id,
|
|
||||||
model_name: stem.to_string(),
|
model_name: stem.to_string(),
|
||||||
repeat_penalty: 1.1,
|
// repeat_penalty: 1.1,
|
||||||
repeat_last_n: 64,
|
// repeat_last_n: 64,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -130,9 +127,6 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
|||||||
let seed = mes.seed.unwrap_or(32768) as u64;
|
let seed = mes.seed.unwrap_or(32768) as u64;
|
||||||
let temperature = mes.temperature.unwrap_or(0.4);
|
let temperature = mes.temperature.unwrap_or(0.4);
|
||||||
let top_p = mes.top_p.unwrap_or(0.95);
|
let top_p = mes.top_p.unwrap_or(0.95);
|
||||||
let mut logit_processor =
|
|
||||||
get_logit_processor(temperature.into(), top_p.into(), Some(20), seed);
|
|
||||||
// let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
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 {
|
||||||
@@ -147,57 +141,32 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
|||||||
} else {
|
} else {
|
||||||
(mes_render, None, None, None, None)
|
(mes_render, None, None, None, None)
|
||||||
};
|
};
|
||||||
let mut input_ids = self.tokenizer.text_encode(mes_text, &self.device)?;
|
let input_ids = self.tokenizer.text_encode(mes_text, &self.device)?;
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
|
||||||
let prompt_tokens = seq_len as u32;
|
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
let mut pixel_values = pixel_values.as_ref();
|
|
||||||
let image_grid_thw = image_grid_thw.as_ref();
|
|
||||||
let mut pixel_values_video = pixel_values_video.as_ref();
|
|
||||||
let video_grid_thw = video_grid_thw.as_ref();
|
|
||||||
let mut generate = Vec::new();
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
for _ in 0..sample_len {
|
let mut ctx = GenerationContext::new(
|
||||||
let logits = self.qwen3_5.forward(
|
temperature.into(),
|
||||||
&input_ids,
|
top_p.into(),
|
||||||
|
Some(20),
|
||||||
|
seed,
|
||||||
|
input_ids.dim(1)?,
|
||||||
|
sample_len,
|
||||||
|
self.device.clone(),
|
||||||
|
);
|
||||||
|
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,
|
||||||
seqlen_offset,
|
];
|
||||||
)?;
|
let data = MultiModalData::new(data_vec);
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
generate_generic(
|
||||||
let logits = if self.repeat_penalty == 1. {
|
&mut self.qwen3_5,
|
||||||
logits
|
&self.tokenizer,
|
||||||
} else {
|
input_ids,
|
||||||
let start_at = generate.len().saturating_sub(self.repeat_last_n);
|
data,
|
||||||
candle_transformers::utils::apply_repeat_penalty(
|
&mut ctx,
|
||||||
&logits,
|
|
||||||
self.repeat_penalty,
|
|
||||||
&generate[start_at..],
|
|
||||||
)?
|
|
||||||
};
|
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
|
||||||
generate.push(next_token);
|
|
||||||
if next_token == self.eos_token_id {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
pixel_values = None;
|
|
||||||
pixel_values_video = None;
|
|
||||||
}
|
|
||||||
let completion_tokens = generate.len() as u32;
|
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
|
||||||
self.qwen3_5.clear_cache();
|
|
||||||
let response = build_completion_response(
|
|
||||||
res,
|
|
||||||
&self.model_name,
|
&self.model_name,
|
||||||
Some(completion_tokens),
|
)
|
||||||
Some(prompt_tokens),
|
|
||||||
);
|
|
||||||
Ok(response)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
@@ -211,10 +180,8 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
|||||||
+ '_,
|
+ '_,
|
||||||
>,
|
>,
|
||||||
> {
|
> {
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
|
||||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
// let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
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)?;
|
||||||
@@ -228,117 +195,30 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
|||||||
} else {
|
} else {
|
||||||
(mes_render, None, None, None, None)
|
(mes_render, None, None, None, None)
|
||||||
};
|
};
|
||||||
let mut input_ids = self.tokenizer.text_encode(mes_text, &self.device)?;
|
let input_ids = self.tokenizer.text_encode(mes_text, &self.device)?;
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
let stream = stream! {
|
let data_vec = vec![
|
||||||
let mut error_tokens = Vec::new();
|
|
||||||
let mut pixel_values = pixel_values.as_ref();
|
|
||||||
let image_grid_thw = image_grid_thw.as_ref();
|
|
||||||
let mut pixel_values_video = pixel_values_video.as_ref();
|
|
||||||
let video_grid_thw = video_grid_thw.as_ref();
|
|
||||||
let mut tool_call_id = None;
|
|
||||||
let mut tool_call_content = String::new();
|
|
||||||
let mut generate = Vec::new();
|
|
||||||
for _ in 0..sample_len {
|
|
||||||
let logits = self.qwen3_5.forward(
|
|
||||||
&input_ids,
|
|
||||||
pixel_values,
|
pixel_values,
|
||||||
image_grid_thw,
|
image_grid_thw,
|
||||||
pixel_values_video,
|
pixel_values_video,
|
||||||
video_grid_thw,
|
video_grid_thw,
|
||||||
seqlen_offset,
|
];
|
||||||
)?;
|
let data = MultiModalData::new(data_vec);
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let logits = if self.repeat_penalty == 1. {
|
let stream = generate_stream_generic(
|
||||||
logits
|
&mut self.qwen3_5,
|
||||||
} else {
|
&self.tokenizer,
|
||||||
let start_at = generate.len().saturating_sub(self.repeat_last_n);
|
input_ids,
|
||||||
candle_transformers::utils::apply_repeat_penalty(
|
data,
|
||||||
&logits,
|
mes.temperature,
|
||||||
self.repeat_penalty,
|
mes.top_p,
|
||||||
&generate[start_at..],
|
|
||||||
)?
|
|
||||||
};
|
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
|
||||||
generate.push(next_token);
|
|
||||||
let mut decode_ids = Vec::new();
|
|
||||||
if !error_tokens.is_empty() {
|
|
||||||
decode_ids.extend_from_slice(&error_tokens);
|
|
||||||
}
|
|
||||||
decode_ids.push(next_token);
|
|
||||||
let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
|
|
||||||
if decoded_token.contains("�") {
|
|
||||||
error_tokens.push(next_token);
|
|
||||||
if error_tokens.len() > 3 {
|
|
||||||
error_tokens.clear();
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
pixel_values = None;
|
|
||||||
pixel_values_video = None;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
error_tokens.clear();
|
|
||||||
// 处理特殊标记和工具调用
|
|
||||||
match decoded_token.as_str() {
|
|
||||||
"<tool_call>" => {
|
|
||||||
// 开始工具调用
|
|
||||||
tool_call_id = Some(uuid::Uuid::new_v4().to_string());
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
pixel_values = None;
|
|
||||||
pixel_values_video = None;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
"</tool_call>" => {
|
|
||||||
// 结束工具调用
|
|
||||||
let chunk = build_completion_chunk_response(
|
|
||||||
decoded_token,
|
|
||||||
&self.model_name,
|
|
||||||
tool_call_id.clone(),
|
|
||||||
Some(tool_call_content.clone())
|
|
||||||
);
|
|
||||||
tool_call_id = None;
|
|
||||||
tool_call_content = String::new();
|
|
||||||
yield Ok(chunk);
|
|
||||||
}
|
|
||||||
_ => {
|
|
||||||
if tool_call_id.is_some() {
|
|
||||||
// 在工具调用过程中,收集工具调用内容
|
|
||||||
tool_call_content.push_str(&decoded_token);
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
pixel_values = None;
|
|
||||||
pixel_values_video = None;
|
|
||||||
continue;
|
|
||||||
} else {
|
|
||||||
// 正常文本输出
|
|
||||||
let chunk = build_completion_chunk_response(
|
|
||||||
decoded_token,
|
|
||||||
&self.model_name,
|
|
||||||
None,
|
None,
|
||||||
None
|
seed,
|
||||||
);
|
sample_len,
|
||||||
yield Ok(chunk);
|
in_reasoning,
|
||||||
}
|
&self.device,
|
||||||
}
|
&self.model_name,
|
||||||
}
|
)?;
|
||||||
if next_token == self.eos_token_id {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
pixel_values = None;
|
|
||||||
pixel_values_video = None;
|
|
||||||
}
|
|
||||||
self.qwen3_5.clear_cache();
|
|
||||||
};
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ use candle_nn::{
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{
|
common::{
|
||||||
|
InferenceModel,
|
||||||
gguf::{GateUpDownMLPGguf, Gguf, ProjKind, QuantizedLinear},
|
gguf::{GateUpDownMLPGguf, Gguf, ProjKind, QuantizedLinear},
|
||||||
modules::{conv1d_depthwise, eager_attention_forward, get_conv1d, softplus},
|
modules::{conv1d_depthwise, eager_attention_forward, get_conv1d, softplus},
|
||||||
},
|
},
|
||||||
@@ -1043,10 +1044,11 @@ pub struct Qwen3_5Model {
|
|||||||
language_model: Qwen3_5TextModel,
|
language_model: Qwen3_5TextModel,
|
||||||
lm_head: ProjKind,
|
lm_head: ProjKind,
|
||||||
rope_deltas: Option<Tensor>,
|
rope_deltas: Option<Tensor>,
|
||||||
|
stop_token_ids: Vec<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Qwen3_5Model {
|
impl Qwen3_5Model {
|
||||||
pub fn new_from_vb(vb: VarBuilder, config: Qwen3_5Config) -> Result<Self> {
|
pub fn new_from_vb(vb: VarBuilder, config: Qwen3_5Config, eos_ids: Vec<u32>) -> Result<Self> {
|
||||||
let vb_m = vb.pp("model");
|
let vb_m = vb.pp("model");
|
||||||
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?;
|
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?;
|
||||||
let language_model =
|
let language_model =
|
||||||
@@ -1069,6 +1071,7 @@ impl Qwen3_5Model {
|
|||||||
language_model,
|
language_model,
|
||||||
lm_head: ProjKind::LinearProj(lm_head),
|
lm_head: ProjKind::LinearProj(lm_head),
|
||||||
rope_deltas: None,
|
rope_deltas: None,
|
||||||
|
stop_token_ids: eos_ids,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1076,6 +1079,7 @@ impl Qwen3_5Model {
|
|||||||
gguf: &mut Gguf<R>,
|
gguf: &mut Gguf<R>,
|
||||||
mmproj_gguf: Option<&mut Gguf<R>>,
|
mmproj_gguf: Option<&mut Gguf<R>>,
|
||||||
device: &Device,
|
device: &Device,
|
||||||
|
eos_ids: Vec<u32>,
|
||||||
) -> Result<Self> {
|
) -> Result<Self> {
|
||||||
let spatial_merge_size = 2usize;
|
let spatial_merge_size = 2usize;
|
||||||
let image_token_id = 248056u32;
|
let image_token_id = 248056u32;
|
||||||
@@ -1102,6 +1106,7 @@ impl Qwen3_5Model {
|
|||||||
language_model,
|
language_model,
|
||||||
lm_head: ProjKind::QuantizedProj(QuantizedLinear::new(lm_head, None)),
|
lm_head: ProjKind::QuantizedProj(QuantizedLinear::new(lm_head, None)),
|
||||||
rope_deltas: None,
|
rope_deltas: None,
|
||||||
|
stop_token_ids: eos_ids,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1434,6 +1439,46 @@ impl Qwen3_5Model {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn clear_cache(&mut self) {
|
pub fn clear_cache(&mut self) {
|
||||||
|
self.rope_deltas = None;
|
||||||
self.language_model.clear_cache();
|
self.language_model.clear_cache();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl InferenceModel for Qwen3_5Model {
|
||||||
|
fn forward_initial(
|
||||||
|
&mut self,
|
||||||
|
input_ids: &Tensor,
|
||||||
|
seqlen_offset: usize,
|
||||||
|
data: crate::models::common::MultiModalData,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
if data.data_vec.len() != 4 {
|
||||||
|
return Err(anyhow::anyhow!(
|
||||||
|
"Lfm2VL process data error, must have pixel_values, image_grid_thw, pixel_values_video, video_grid_thw"
|
||||||
|
));
|
||||||
|
}
|
||||||
|
let pixel_values = &data.data_vec[0];
|
||||||
|
let image_grid_thw = &data.data_vec[1];
|
||||||
|
let pixel_values_video = &data.data_vec[2];
|
||||||
|
let video_grid_thw = &data.data_vec[3];
|
||||||
|
self.forward(
|
||||||
|
input_ids,
|
||||||
|
pixel_values.as_ref(),
|
||||||
|
image_grid_thw.as_ref(),
|
||||||
|
pixel_values_video.as_ref(),
|
||||||
|
video_grid_thw.as_ref(),
|
||||||
|
seqlen_offset,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
|
self.forward(input_ids, None, None, None, None, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clear_cache(&mut self) {
|
||||||
|
self.clear_cache();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn stop_token_ids(&self) -> Vec<u32> {
|
||||||
|
self.stop_token_ids.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -202,7 +202,7 @@ pub struct Qwen3ASRRopeScaling {
|
|||||||
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||||
pub struct Qwen3ASRGenerationConfig {
|
pub struct Qwen3ASRGenerationConfig {
|
||||||
pub do_sample: bool,
|
pub do_sample: bool,
|
||||||
pub eos_token_id: Vec<usize>,
|
pub eos_token_id: Vec<u32>,
|
||||||
pub pad_token_id: usize,
|
pub pad_token_id: usize,
|
||||||
pub temperature: f32,
|
pub temperature: f32,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,6 +1,9 @@
|
|||||||
|
use std::time::Instant;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
models::common::generate::get_logit_processor,
|
models::common::generate::get_logit_processor,
|
||||||
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
|
utils::response_utils::{build_chunk_response_with_usage, build_completion_response_with_time},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device, Tensor};
|
||||||
@@ -21,8 +24,7 @@ use crate::{
|
|||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
find_type_files, get_device, get_dtype, response_utils::build_completion_chunk_response,
|
||||||
get_dtype,
|
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -92,6 +94,8 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
|
|||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
let mut generate = Vec::new();
|
let mut generate = Vec::new();
|
||||||
let mut prompt_tokens = 0u32;
|
let mut prompt_tokens = 0u32;
|
||||||
|
let mut prompt_secs = 0.0f64;
|
||||||
|
let mut completion_secs = 0.0f64;
|
||||||
for data in audio_datas.iter() {
|
for data in audio_datas.iter() {
|
||||||
let mut input_ids = data.input_ids.clone();
|
let mut input_ids = data.input_ids.clone();
|
||||||
let mut input_features = Some(data.input_features.clone().to_dtype(self.dtype)?);
|
let mut input_features = Some(data.input_features.clone().to_dtype(self.dtype)?);
|
||||||
@@ -99,11 +103,18 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
|
|||||||
prompt_tokens += seq_len as u32;
|
prompt_tokens += seq_len as u32;
|
||||||
let mut seqlen_offset = 0;
|
let mut seqlen_offset = 0;
|
||||||
for _ in 0..sample_len {
|
for _ in 0..sample_len {
|
||||||
|
let i_start = Instant::now();
|
||||||
let logits =
|
let logits =
|
||||||
self.qwen3_asr
|
self.qwen3_asr
|
||||||
.forward(&input_ids, seqlen_offset, input_features.as_ref())?;
|
.forward(&input_ids, seqlen_offset, input_features.as_ref())?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
let next_token = logit_processor.sample(&logits)?;
|
||||||
|
let i_duration = i_start.elapsed();
|
||||||
|
if seqlen_offset == 0 {
|
||||||
|
prompt_secs += i_duration.as_secs_f64();
|
||||||
|
} else {
|
||||||
|
completion_secs += i_duration.as_secs_f64();
|
||||||
|
};
|
||||||
generate.push(next_token);
|
generate.push(next_token);
|
||||||
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
||||||
break;
|
break;
|
||||||
@@ -117,8 +128,14 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
|
|||||||
}
|
}
|
||||||
let num_token = generate.len() as u32;
|
let num_token = generate.len() as u32;
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
let res = self.tokenizer.token_decode(generate)?;
|
||||||
let response =
|
let response = build_completion_response_with_time(
|
||||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
res,
|
||||||
|
&self.model_name,
|
||||||
|
num_token.into(),
|
||||||
|
completion_secs.into(),
|
||||||
|
prompt_tokens.into(),
|
||||||
|
prompt_secs.into(),
|
||||||
|
);
|
||||||
Ok(response)
|
Ok(response)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -145,17 +162,30 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
|
|||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
let stream = stream! {
|
let stream = stream! {
|
||||||
let mut error_tokens = Vec::new();
|
let mut error_tokens = Vec::new();
|
||||||
|
let mut prompt_tokens = 0u32;
|
||||||
|
let mut completion_tokens = 0u32;
|
||||||
|
let mut prompt_secs = 0.0f64;
|
||||||
|
let mut completion_secs = 0.0f64;
|
||||||
for data in audio_datas.iter() {
|
for data in audio_datas.iter() {
|
||||||
let mut input_ids = data.input_ids.clone();
|
let mut input_ids = data.input_ids.clone();
|
||||||
let mut input_features = Some(data.input_features.clone().to_dtype(self.dtype)?);
|
let mut input_features = Some(data.input_features.clone().to_dtype(self.dtype)?);
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
let mut seq_len = input_ids.dim(1)?;
|
||||||
|
prompt_tokens += seq_len as u32;
|
||||||
let mut seqlen_offset = 0;
|
let mut seqlen_offset = 0;
|
||||||
for _ in 0..sample_len {
|
for _ in 0..sample_len {
|
||||||
|
let i_start = Instant::now();
|
||||||
let logits =
|
let logits =
|
||||||
self.qwen3_asr
|
self.qwen3_asr
|
||||||
.forward(&input_ids, seqlen_offset, input_features.as_ref())?;
|
.forward(&input_ids, seqlen_offset, input_features.as_ref())?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
let next_token = logit_processor.sample(&logits)?;
|
||||||
|
completion_tokens += 1;
|
||||||
|
let i_duration = i_start.elapsed();
|
||||||
|
if seqlen_offset == 0 {
|
||||||
|
prompt_secs += i_duration.as_secs_f64();
|
||||||
|
} else {
|
||||||
|
completion_secs += i_duration.as_secs_f64();
|
||||||
|
};
|
||||||
let mut decode_ids = Vec::new();
|
let mut decode_ids = Vec::new();
|
||||||
if !error_tokens.is_empty() {
|
if !error_tokens.is_empty() {
|
||||||
decode_ids.extend_from_slice(&error_tokens);
|
decode_ids.extend_from_slice(&error_tokens);
|
||||||
@@ -177,6 +207,7 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
|
|||||||
let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
|
||||||
yield Ok(chunk);
|
yield Ok(chunk);
|
||||||
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
||||||
|
yield Ok(build_chunk_response_with_usage(&self.model_name, completion_tokens.into(), completion_secs.into(), prompt_tokens.into(), prompt_secs.into()));
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
seqlen_offset += seq_len;
|
seqlen_offset += seq_len;
|
||||||
|
|||||||
+165
-164
@@ -1,11 +1,13 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::common::generate::get_logit_processor,
|
models::common::{
|
||||||
|
MultiModalData,
|
||||||
|
generate::{GenerationContext, generate_generic, generate_stream_generic},
|
||||||
|
},
|
||||||
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
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::async_stream::stream;
|
|
||||||
use rocket::futures::Stream;
|
use rocket::futures::Stream;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
@@ -16,10 +18,7 @@ use crate::{
|
|||||||
qwen3vl::{config::Qwen3VLConfig, model::Qwen3VLModel, processor::Qwen3VLProcessor},
|
qwen3vl::{config::Qwen3VLConfig, model::Qwen3VLModel, processor::Qwen3VLProcessor},
|
||||||
},
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{find_type_files, get_device, get_dtype},
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
|
||||||
get_dtype,
|
|
||||||
},
|
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct Qwen3VLGenerateModel<'a> {
|
pub struct Qwen3VLGenerateModel<'a> {
|
||||||
@@ -28,8 +27,6 @@ pub struct Qwen3VLGenerateModel<'a> {
|
|||||||
pre_processor: Qwen3VLProcessor,
|
pre_processor: Qwen3VLProcessor,
|
||||||
qwen3_vl: Qwen3VLModel,
|
qwen3_vl: Qwen3VLModel,
|
||||||
device: Device,
|
device: Device,
|
||||||
eos_token_id1: u32,
|
|
||||||
eos_token_id2: u32,
|
|
||||||
generation_config: Qwen3GenerationConfig,
|
generation_config: Qwen3GenerationConfig,
|
||||||
model_name: String,
|
model_name: String,
|
||||||
}
|
}
|
||||||
@@ -46,10 +43,11 @@ impl<'a> Qwen3VLGenerateModel<'a> {
|
|||||||
let pre_processor = Qwen3VLProcessor::new(path, &device, dtype)?;
|
let pre_processor = Qwen3VLProcessor::new(path, &device, dtype)?;
|
||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
||||||
let qwen3_vl = Qwen3VLModel::new(cfg, vb)?;
|
|
||||||
let generation_config_path = path.to_string() + "/generation_config.json";
|
let generation_config_path = path.to_string() + "/generation_config.json";
|
||||||
let generation_config: Qwen3GenerationConfig =
|
let generation_config: Qwen3GenerationConfig =
|
||||||
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
|
||||||
|
let qwen3_vl = Qwen3VLModel::new(cfg, vb, generation_config.eos_token_id.clone())?;
|
||||||
|
|
||||||
let model_name = std::path::Path::new(path)
|
let model_name = std::path::Path::new(path)
|
||||||
.file_name()
|
.file_name()
|
||||||
.and_then(|s| s.to_str())
|
.and_then(|s| s.to_str())
|
||||||
@@ -61,8 +59,6 @@ impl<'a> Qwen3VLGenerateModel<'a> {
|
|||||||
pre_processor,
|
pre_processor,
|
||||||
qwen3_vl,
|
qwen3_vl,
|
||||||
device,
|
device,
|
||||||
eos_token_id1: generation_config.eos_token_id[0] as u32,
|
|
||||||
eos_token_id2: generation_config.eos_token_id[1] as u32,
|
|
||||||
generation_config,
|
generation_config,
|
||||||
model_name,
|
model_name,
|
||||||
})
|
})
|
||||||
@@ -77,56 +73,39 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
|||||||
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
||||||
let top_k = self.generation_config.top_k;
|
let top_k = self.generation_config.top_k;
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let mut logit_processor =
|
|
||||||
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
// let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
|
|
||||||
// let mes_render = self
|
|
||||||
// .chat_template
|
|
||||||
// .apply_chat_temp_think(&mes, enable_thinking)?;
|
|
||||||
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||||
let mut 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 mut seq_len = input_ids.dim(1)?;
|
let seq_len = input_ids.dim(1)?;
|
||||||
let prompt_tokens = seq_len as u32;
|
|
||||||
let mut seqlen_offset = 0;
|
|
||||||
let mut pixel_values = input.pixel_values.as_ref();
|
|
||||||
let image_grid_thw = input.image_grid_thw.as_ref();
|
|
||||||
let mut pixel_values_video = input.pixel_values_video.as_ref();
|
|
||||||
let video_grid_thw = input.video_grid_thw.as_ref();
|
|
||||||
let mut cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?;
|
|
||||||
let mut generate = Vec::new();
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
for _ in 0..sample_len {
|
let mut ctx = GenerationContext::new(
|
||||||
let logits = self.qwen3_vl.forward(
|
temperature.into(),
|
||||||
&input_ids,
|
top_p.into(),
|
||||||
pixel_values,
|
top_k.into(),
|
||||||
image_grid_thw,
|
seed,
|
||||||
pixel_values_video,
|
input_ids.dim(1)?,
|
||||||
video_grid_thw,
|
sample_len,
|
||||||
Some(&cache_position),
|
self.device.clone(),
|
||||||
seqlen_offset,
|
);
|
||||||
)?;
|
let cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
let data_vec = vec![
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
input.pixel_values,
|
||||||
generate.push(next_token);
|
input.image_grid_thw,
|
||||||
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
input.pixel_values_video,
|
||||||
break;
|
input.video_grid_thw,
|
||||||
}
|
cache_position.into(),
|
||||||
seqlen_offset += seq_len;
|
];
|
||||||
seq_len = 1;
|
let data = MultiModalData::new(data_vec);
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
generate_generic(
|
||||||
cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
&mut self.qwen3_vl,
|
||||||
pixel_values = None;
|
&self.tokenizer,
|
||||||
pixel_values_video = None;
|
input_ids,
|
||||||
}
|
data,
|
||||||
let num_token = generate.len() as u32;
|
&mut ctx,
|
||||||
let res = self.tokenizer.token_decode(generate)?;
|
&self.model_name,
|
||||||
self.qwen3_vl.clear_kv_cache();
|
)
|
||||||
let response =
|
|
||||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
|
||||||
Ok(response)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
@@ -145,119 +124,141 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
|||||||
.unwrap_or(self.generation_config.temperature);
|
.unwrap_or(self.generation_config.temperature);
|
||||||
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
|
||||||
let top_k = self.generation_config.top_k;
|
let top_k = self.generation_config.top_k;
|
||||||
let seed = mes.seed.unwrap_or(34562) as u64;
|
|
||||||
let mut logit_processor =
|
|
||||||
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
|
||||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
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 = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||||
let mut 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 mut seq_len = input_ids.dim(1)?;
|
let seq_len = input_ids.dim(1)?;
|
||||||
let mut seqlen_offset = 0;
|
let cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?;
|
||||||
let mut cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?;
|
|
||||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
let stream = stream! {
|
let data_vec = vec![
|
||||||
let mut error_tokens = Vec::new();
|
input.pixel_values,
|
||||||
let mut pixel_values = input.pixel_values.as_ref();
|
input.image_grid_thw,
|
||||||
let image_grid_thw = input.image_grid_thw.as_ref();
|
input.pixel_values_video,
|
||||||
let mut pixel_values_video = input.pixel_values_video.as_ref();
|
input.video_grid_thw,
|
||||||
let video_grid_thw = input.video_grid_thw.as_ref();
|
cache_position.into(),
|
||||||
let mut tool_call_id = None;
|
];
|
||||||
let mut tool_call_content = String::new();
|
let data = MultiModalData::new(data_vec);
|
||||||
for _ in 0..sample_len {
|
let seed = mes.seed.unwrap_or(34562) as u64;
|
||||||
let logits = self.qwen3_vl.forward(
|
let stream = generate_stream_generic(
|
||||||
&input_ids,
|
&mut self.qwen3_vl,
|
||||||
pixel_values,
|
&self.tokenizer,
|
||||||
image_grid_thw,
|
input_ids,
|
||||||
pixel_values_video,
|
data,
|
||||||
video_grid_thw,
|
temperature.into(),
|
||||||
Some(&cache_position),
|
top_p.into(),
|
||||||
seqlen_offset,
|
top_k.into(),
|
||||||
|
seed,
|
||||||
|
sample_len,
|
||||||
|
in_reasoning,
|
||||||
|
&self.device,
|
||||||
|
&self.model_name,
|
||||||
)?;
|
)?;
|
||||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
|
||||||
let next_token = logit_processor.sample(&logits)?;
|
|
||||||
let mut decode_ids = Vec::new();
|
|
||||||
if !error_tokens.is_empty() {
|
|
||||||
decode_ids.extend_from_slice(&error_tokens);
|
|
||||||
}
|
|
||||||
decode_ids.push(next_token);
|
|
||||||
let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
|
|
||||||
if decoded_token.contains("�") {
|
|
||||||
error_tokens.push(next_token);
|
|
||||||
if error_tokens.len() > 3 {
|
|
||||||
error_tokens.clear();
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
|
||||||
pixel_values = None;
|
|
||||||
pixel_values_video = None;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
error_tokens.clear();
|
|
||||||
|
|
||||||
// 处理特殊标记和工具调用
|
|
||||||
match decoded_token.as_str() {
|
|
||||||
"<tool_call>" => {
|
|
||||||
// 开始工具调用
|
|
||||||
tool_call_id = Some(uuid::Uuid::new_v4().to_string());
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
|
||||||
pixel_values = None;
|
|
||||||
pixel_values_video = None;
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
"</tool_call>" => {
|
|
||||||
// 结束工具调用
|
|
||||||
let chunk = build_completion_chunk_response(
|
|
||||||
decoded_token,
|
|
||||||
&self.model_name,
|
|
||||||
tool_call_id.clone(),
|
|
||||||
Some(tool_call_content.clone())
|
|
||||||
);
|
|
||||||
tool_call_id = None;
|
|
||||||
tool_call_content = String::new();
|
|
||||||
yield Ok(chunk);
|
|
||||||
}
|
|
||||||
_ => {
|
|
||||||
if tool_call_id.is_some() {
|
|
||||||
// 在工具调用过程中,收集工具调用内容
|
|
||||||
tool_call_content.push_str(&decoded_token);
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
|
||||||
pixel_values = None;
|
|
||||||
pixel_values_video = None;
|
|
||||||
continue;
|
|
||||||
} else {
|
|
||||||
// 正常文本输出
|
|
||||||
let chunk = build_completion_chunk_response(
|
|
||||||
decoded_token,
|
|
||||||
&self.model_name,
|
|
||||||
None,
|
|
||||||
None
|
|
||||||
);
|
|
||||||
yield Ok(chunk);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
seqlen_offset += seq_len;
|
|
||||||
seq_len = 1;
|
|
||||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
|
||||||
cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
|
||||||
pixel_values = None;
|
|
||||||
pixel_values_video = None;
|
|
||||||
}
|
|
||||||
self.qwen3_vl.clear_kv_cache();
|
|
||||||
};
|
|
||||||
Ok(Box::new(Box::pin(stream)))
|
Ok(Box::new(Box::pin(stream)))
|
||||||
|
// let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||||
|
// let stream = stream! {
|
||||||
|
// let mut error_tokens = Vec::new();
|
||||||
|
// let mut pixel_values = input.pixel_values.as_ref();
|
||||||
|
// let image_grid_thw = input.image_grid_thw.as_ref();
|
||||||
|
// let mut pixel_values_video = input.pixel_values_video.as_ref();
|
||||||
|
// let video_grid_thw = input.video_grid_thw.as_ref();
|
||||||
|
// let mut tool_call_id = None;
|
||||||
|
// let mut tool_call_content = String::new();
|
||||||
|
// for _ in 0..sample_len {
|
||||||
|
// let logits = self.qwen3_vl.forward(
|
||||||
|
// &input_ids,
|
||||||
|
// pixel_values,
|
||||||
|
// image_grid_thw,
|
||||||
|
// pixel_values_video,
|
||||||
|
// video_grid_thw,
|
||||||
|
// Some(&cache_position),
|
||||||
|
// seqlen_offset,
|
||||||
|
// )?;
|
||||||
|
// let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||||
|
// let next_token = logit_processor.sample(&logits)?;
|
||||||
|
// let mut decode_ids = Vec::new();
|
||||||
|
// if !error_tokens.is_empty() {
|
||||||
|
// decode_ids.extend_from_slice(&error_tokens);
|
||||||
|
// }
|
||||||
|
// decode_ids.push(next_token);
|
||||||
|
// let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
|
||||||
|
// if decoded_token.contains("�") {
|
||||||
|
// error_tokens.push(next_token);
|
||||||
|
// if error_tokens.len() > 3 {
|
||||||
|
// error_tokens.clear();
|
||||||
|
// }
|
||||||
|
// seqlen_offset += seq_len;
|
||||||
|
// seq_len = 1;
|
||||||
|
// input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||||
|
// cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
||||||
|
// pixel_values = None;
|
||||||
|
// pixel_values_video = None;
|
||||||
|
// continue;
|
||||||
|
// }
|
||||||
|
// error_tokens.clear();
|
||||||
|
|
||||||
|
// // 处理特殊标记和工具调用
|
||||||
|
// match decoded_token.as_str() {
|
||||||
|
// "<tool_call>" => {
|
||||||
|
// // 开始工具调用
|
||||||
|
// tool_call_id = Some(uuid::Uuid::new_v4().to_string());
|
||||||
|
// seqlen_offset += seq_len;
|
||||||
|
// seq_len = 1;
|
||||||
|
// input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||||
|
// cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
||||||
|
// pixel_values = None;
|
||||||
|
// pixel_values_video = None;
|
||||||
|
// continue;
|
||||||
|
// }
|
||||||
|
// "</tool_call>" => {
|
||||||
|
// // 结束工具调用
|
||||||
|
// let chunk = build_completion_chunk_response(
|
||||||
|
// decoded_token,
|
||||||
|
// &self.model_name,
|
||||||
|
// tool_call_id.clone(),
|
||||||
|
// Some(tool_call_content.clone())
|
||||||
|
// );
|
||||||
|
// tool_call_id = None;
|
||||||
|
// tool_call_content = String::new();
|
||||||
|
// yield Ok(chunk);
|
||||||
|
// }
|
||||||
|
// _ => {
|
||||||
|
// if tool_call_id.is_some() {
|
||||||
|
// // 在工具调用过程中,收集工具调用内容
|
||||||
|
// tool_call_content.push_str(&decoded_token);
|
||||||
|
// seqlen_offset += seq_len;
|
||||||
|
// seq_len = 1;
|
||||||
|
// input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||||
|
// cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
||||||
|
// pixel_values = None;
|
||||||
|
// pixel_values_video = None;
|
||||||
|
// continue;
|
||||||
|
// } else {
|
||||||
|
// // 正常文本输出
|
||||||
|
// let chunk = build_completion_chunk_response(
|
||||||
|
// decoded_token,
|
||||||
|
// &self.model_name,
|
||||||
|
// None,
|
||||||
|
// None
|
||||||
|
// );
|
||||||
|
// yield Ok(chunk);
|
||||||
|
// }
|
||||||
|
// }
|
||||||
|
// }
|
||||||
|
// if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
|
||||||
|
// break;
|
||||||
|
// }
|
||||||
|
// seqlen_offset += seq_len;
|
||||||
|
// seq_len = 1;
|
||||||
|
// input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||||
|
// cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
|
||||||
|
// pixel_values = None;
|
||||||
|
// pixel_values_video = None;
|
||||||
|
// }
|
||||||
|
// self.qwen3_vl.clear_kv_cache();
|
||||||
|
// };
|
||||||
|
// Ok(Box::new(Box::pin(stream)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ use candle_nn::{
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::{
|
models::{
|
||||||
common::{
|
common::{
|
||||||
|
InferenceModel,
|
||||||
gguf::{Gguf, ProjKind, TwoLinearMLPGguf},
|
gguf::{Gguf, ProjKind, TwoLinearMLPGguf},
|
||||||
modules::{eager_attention_forward, get_layer_norm},
|
modules::{eager_attention_forward, get_layer_norm},
|
||||||
},
|
},
|
||||||
@@ -839,10 +840,11 @@ pub struct Qwen3VLModel {
|
|||||||
language_model: Qwen3VLTextModel,
|
language_model: Qwen3VLTextModel,
|
||||||
lm_head: Linear,
|
lm_head: Linear,
|
||||||
rope_deltas: Option<Tensor>,
|
rope_deltas: Option<Tensor>,
|
||||||
|
stop_token_ids: Vec<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Qwen3VLModel {
|
impl Qwen3VLModel {
|
||||||
pub fn new(config: Qwen3VLConfig, vb: VarBuilder) -> Result<Self> {
|
pub fn new(config: Qwen3VLConfig, vb: VarBuilder, eos_ids: Vec<u32>) -> Result<Self> {
|
||||||
let vb_m = vb.pp("model");
|
let vb_m = vb.pp("model");
|
||||||
let config = config.clone();
|
let config = config.clone();
|
||||||
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?;
|
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?;
|
||||||
@@ -863,6 +865,7 @@ impl Qwen3VLModel {
|
|||||||
language_model,
|
language_model,
|
||||||
lm_head,
|
lm_head,
|
||||||
rope_deltas: None,
|
rope_deltas: None,
|
||||||
|
stop_token_ids: eos_ids,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1240,6 +1243,15 @@ impl Qwen3VLModel {
|
|||||||
.broadcast_add(rope_deltas)?
|
.broadcast_add(rope_deltas)?
|
||||||
.contiguous()?
|
.contiguous()?
|
||||||
.to_dtype(candle_core::DType::U32)?
|
.to_dtype(candle_core::DType::U32)?
|
||||||
|
} else if let Some(rope_deltas) = &self.rope_deltas {
|
||||||
|
let cache_position =
|
||||||
|
Tensor::from_vec(vec![seqlen_offset as u32], 1, inputs_embeds.device())?;
|
||||||
|
cache_position
|
||||||
|
.i(0)?
|
||||||
|
.to_dtype(rope_deltas.dtype())?
|
||||||
|
.broadcast_add(rope_deltas)?
|
||||||
|
.contiguous()?
|
||||||
|
.to_dtype(candle_core::DType::U32)?
|
||||||
} else {
|
} else {
|
||||||
Tensor::zeros(1, inputs_embeds.dtype(), inputs_embeds.device())?
|
Tensor::zeros(1, inputs_embeds.dtype(), inputs_embeds.device())?
|
||||||
};
|
};
|
||||||
@@ -1265,6 +1277,48 @@ impl Qwen3VLModel {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn clear_kv_cache(&mut self) {
|
pub fn clear_kv_cache(&mut self) {
|
||||||
|
self.rope_deltas = None;
|
||||||
self.language_model.clear_kv_cache();
|
self.language_model.clear_kv_cache();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl InferenceModel for Qwen3VLModel {
|
||||||
|
fn forward_initial(
|
||||||
|
&mut self,
|
||||||
|
input_ids: &Tensor,
|
||||||
|
seqlen_offset: usize,
|
||||||
|
data: crate::models::common::MultiModalData,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
if data.data_vec.len() != 5 {
|
||||||
|
return Err(anyhow::anyhow!(
|
||||||
|
"Qwen3VL process data error, must have pixel_values, image_grid_thw, pixel_values_video, video_grid_thw, cache_position"
|
||||||
|
));
|
||||||
|
}
|
||||||
|
let pixel_values = &data.data_vec[0];
|
||||||
|
let image_grid_thw = &data.data_vec[1];
|
||||||
|
let pixel_values_video = &data.data_vec[2];
|
||||||
|
let video_grid_thw = &data.data_vec[3];
|
||||||
|
let cache_position = &data.data_vec[4];
|
||||||
|
self.forward(
|
||||||
|
input_ids,
|
||||||
|
pixel_values.as_ref(),
|
||||||
|
image_grid_thw.as_ref(),
|
||||||
|
pixel_values_video.as_ref(),
|
||||||
|
video_grid_thw.as_ref(),
|
||||||
|
cache_position.as_ref(),
|
||||||
|
seqlen_offset,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
|
||||||
|
self.forward(input_ids, None, None, None, None, None, seqlen_offset)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clear_cache(&mut self) {
|
||||||
|
self.clear_kv_cache();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn stop_token_ids(&self) -> Vec<u32> {
|
||||||
|
self.stop_token_ids.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -14,8 +14,9 @@ use rocket::futures::{Stream, stream};
|
|||||||
use crate::{
|
use crate::{
|
||||||
models::{GenerateModel, rmbg2_0::model::BiRefNet},
|
models::{GenerateModel, rmbg2_0::model::BiRefNet},
|
||||||
utils::{
|
utils::{
|
||||||
build_img_completion_response, find_type_files, get_device, get_dtype,
|
find_type_files, get_device, get_dtype,
|
||||||
img_utils::{extract_images, float_tensor_to_dynamic_image, img_transform_with_resize},
|
img_utils::{extract_images, float_tensor_to_dynamic_image, img_transform_with_resize},
|
||||||
|
response_utils::build_img_completion_response,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -21,8 +21,8 @@ use crate::{
|
|||||||
},
|
},
|
||||||
utils::{
|
utils::{
|
||||||
audio_utils::{extract_audio_url, get_audio_wav_u8},
|
audio_utils::{extract_audio_url, get_audio_wav_u8},
|
||||||
build_audio_completion_response, extract_metadata_value, extract_user_text,
|
extract_metadata_value, extract_user_text, find_type_files, get_device, get_dtype,
|
||||||
find_type_files, get_device, get_dtype,
|
response_utils::build_audio_completion_response,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -70,6 +70,8 @@ pub struct ChatCompletionParameters {
|
|||||||
/// Developer-defined tags and values used for filtering completions in the dashboard.
|
/// Developer-defined tags and values used for filtering completions in the dashboard.
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
pub metadata: Option<HashMap<String, String>>,
|
pub metadata: Option<HashMap<String, String>>,
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
pub enable_thinking: Option<bool>,
|
||||||
/// Number between -2.0 and 2.0. Positive values penalize new tokens based on their existing frequency in the text so far,
|
/// Number between -2.0 and 2.0. Positive values penalize new tokens based on their existing frequency in the text so far,
|
||||||
/// decreasing the model's likelihood to repeat the same line verbatim.
|
/// decreasing the model's likelihood to repeat the same line verbatim.
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
|||||||
+3
-307
@@ -1,6 +1,7 @@
|
|||||||
pub mod audio_utils;
|
pub mod audio_utils;
|
||||||
pub mod img_utils;
|
pub mod img_utils;
|
||||||
pub mod interpolate;
|
pub mod interpolate;
|
||||||
|
pub mod response_utils;
|
||||||
pub mod tensor_utils;
|
pub mod tensor_utils;
|
||||||
pub mod video_utils;
|
pub mod video_utils;
|
||||||
|
|
||||||
@@ -10,14 +11,8 @@ use std::time::{SystemTime, UNIX_EPOCH};
|
|||||||
use std::{collections::HashMap, fs, path::PathBuf, process::Command, time::Duration};
|
use std::{collections::HashMap, fs, path::PathBuf, process::Command, time::Duration};
|
||||||
|
|
||||||
use crate::models::common::model_mapping::WhichModel;
|
use crate::models::common::model_mapping::WhichModel;
|
||||||
use crate::params::{
|
use crate::params::chat::{
|
||||||
chat::{
|
ChatCompletionParameters, ChatMessage, ChatMessageContent, ChatMessageContentPart,
|
||||||
AudioUrlType, ChatCompletionChoice, ChatCompletionChunkChoice, ChatCompletionChunkResponse,
|
|
||||||
ChatCompletionParameters, ChatCompletionResponse, ChatMessage, ChatMessageAudioContentPart,
|
|
||||||
ChatMessageContent, ChatMessageContentPart, ChatMessageImageContentPart, DeltaChatMessage,
|
|
||||||
DeltaFunction, DeltaToolCall, Function, ImageUrlType, ToolCall,
|
|
||||||
},
|
|
||||||
shared::{FinishReason, Usage},
|
|
||||||
};
|
};
|
||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use byteorder::{LittleEndian, ReadBytesExt};
|
use byteorder::{LittleEndian, ReadBytesExt};
|
||||||
@@ -409,305 +404,6 @@ pub fn ceil_by_factor(num: f32, factor: u32) -> u32 {
|
|||||||
ceil * factor
|
ceil * factor
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn build_img_completion_response(
|
|
||||||
base64vec: &Vec<String>,
|
|
||||||
model_name: &str,
|
|
||||||
) -> ChatCompletionResponse {
|
|
||||||
let id = uuid::Uuid::new_v4().to_string();
|
|
||||||
let mut response = ChatCompletionResponse {
|
|
||||||
id: Some(id),
|
|
||||||
choices: vec![],
|
|
||||||
// created: chrono::Utc::now().timestamp() as u32,
|
|
||||||
created: timestamp() as u32,
|
|
||||||
model: model_name.to_string(),
|
|
||||||
service_tier: None,
|
|
||||||
system_fingerprint: None,
|
|
||||||
object: "chat.completion".to_string(),
|
|
||||||
usage: None,
|
|
||||||
};
|
|
||||||
let mut conten_part_vec = vec![];
|
|
||||||
for img_bas64 in base64vec {
|
|
||||||
let img_base64_prefix = "data:image/png;base64,".to_string() + img_bas64;
|
|
||||||
let part = ChatMessageContentPart::Image(ChatMessageImageContentPart {
|
|
||||||
r#type: "image".to_string(),
|
|
||||||
image_url: ImageUrlType {
|
|
||||||
url: img_base64_prefix,
|
|
||||||
detail: None,
|
|
||||||
},
|
|
||||||
});
|
|
||||||
conten_part_vec.push(part);
|
|
||||||
}
|
|
||||||
let choice = ChatCompletionChoice {
|
|
||||||
index: 0,
|
|
||||||
message: ChatMessage::Assistant {
|
|
||||||
content: Some(ChatMessageContent::ContentPart(conten_part_vec)),
|
|
||||||
reasoning_content: None,
|
|
||||||
refusal: None,
|
|
||||||
name: None,
|
|
||||||
audio: None,
|
|
||||||
tool_calls: None,
|
|
||||||
},
|
|
||||||
finish_reason: Some(FinishReason::StopSequenceReached),
|
|
||||||
logprobs: None,
|
|
||||||
};
|
|
||||||
response.choices.push(choice);
|
|
||||||
response
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn build_audio_completion_response(
|
|
||||||
base64_audio: &String,
|
|
||||||
model_name: &str,
|
|
||||||
) -> ChatCompletionResponse {
|
|
||||||
let id = uuid::Uuid::new_v4().to_string();
|
|
||||||
let mut response = ChatCompletionResponse {
|
|
||||||
id: Some(id),
|
|
||||||
choices: vec![],
|
|
||||||
created: timestamp() as u32,
|
|
||||||
model: model_name.to_string(),
|
|
||||||
service_tier: None,
|
|
||||||
system_fingerprint: None,
|
|
||||||
object: "chat.completion".to_string(),
|
|
||||||
usage: None,
|
|
||||||
};
|
|
||||||
|
|
||||||
let base64_audio = format!("data:audio/wav;base64,{}", base64_audio);
|
|
||||||
let conten_part_vec = vec![ChatMessageContentPart::Audio(ChatMessageAudioContentPart {
|
|
||||||
r#type: "audio".to_string(),
|
|
||||||
audio_url: AudioUrlType {
|
|
||||||
url: base64_audio.to_string(),
|
|
||||||
},
|
|
||||||
})];
|
|
||||||
let choice = ChatCompletionChoice {
|
|
||||||
index: 0,
|
|
||||||
message: ChatMessage::Assistant {
|
|
||||||
content: Some(ChatMessageContent::ContentPart(conten_part_vec)),
|
|
||||||
reasoning_content: None,
|
|
||||||
refusal: None,
|
|
||||||
name: None,
|
|
||||||
audio: None,
|
|
||||||
tool_calls: None,
|
|
||||||
},
|
|
||||||
finish_reason: Some(FinishReason::StopSequenceReached),
|
|
||||||
logprobs: None,
|
|
||||||
};
|
|
||||||
response.choices.push(choice);
|
|
||||||
response
|
|
||||||
}
|
|
||||||
|
|
||||||
fn build_response(res: String, model_name: &str, usage: Option<Usage>) -> ChatCompletionResponse {
|
|
||||||
let id = uuid::Uuid::new_v4().to_string();
|
|
||||||
let mut response = ChatCompletionResponse {
|
|
||||||
id: Some(id),
|
|
||||||
choices: vec![],
|
|
||||||
created: timestamp() as u32,
|
|
||||||
model: model_name.to_string(),
|
|
||||||
service_tier: None,
|
|
||||||
system_fingerprint: None,
|
|
||||||
object: "chat.completion".to_string(),
|
|
||||||
usage,
|
|
||||||
};
|
|
||||||
let choice = if res.contains("<tool_call>") {
|
|
||||||
let mes: Vec<&str> = res.split("<tool_call>").collect();
|
|
||||||
let content = mes[0].to_string();
|
|
||||||
let mut tool_vec = Vec::new();
|
|
||||||
for (i, m) in mes.iter().enumerate().skip(1) {
|
|
||||||
let tool_mes = m.replace("</tool_call>", "");
|
|
||||||
let function = match serde_json::from_str::<serde_json::Value>(&tool_mes) {
|
|
||||||
Ok(json_value) => {
|
|
||||||
let name = json_value
|
|
||||||
.get("name")
|
|
||||||
.and_then(|v| v.as_str())
|
|
||||||
.map(|s| s.to_string())
|
|
||||||
.unwrap_or_default();
|
|
||||||
|
|
||||||
let arguments = json_value
|
|
||||||
.get("arguments")
|
|
||||||
.map(|v| v.to_string())
|
|
||||||
.unwrap_or_default();
|
|
||||||
|
|
||||||
Function { name, arguments }
|
|
||||||
}
|
|
||||||
Err(_) => Function {
|
|
||||||
name: "".to_string(),
|
|
||||||
arguments: "".to_string(),
|
|
||||||
},
|
|
||||||
};
|
|
||||||
let tool_call = ToolCall {
|
|
||||||
id: (i - 1).to_string(),
|
|
||||||
r#type: "function".to_string(),
|
|
||||||
function,
|
|
||||||
};
|
|
||||||
tool_vec.push(tool_call);
|
|
||||||
}
|
|
||||||
ChatCompletionChoice {
|
|
||||||
index: 0,
|
|
||||||
message: ChatMessage::Assistant {
|
|
||||||
content: Some(ChatMessageContent::Text(content)),
|
|
||||||
reasoning_content: None,
|
|
||||||
refusal: None,
|
|
||||||
name: None,
|
|
||||||
audio: None,
|
|
||||||
tool_calls: Some(tool_vec),
|
|
||||||
},
|
|
||||||
finish_reason: Some(FinishReason::ToolCalls),
|
|
||||||
logprobs: None,
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
ChatCompletionChoice {
|
|
||||||
index: 0,
|
|
||||||
message: ChatMessage::Assistant {
|
|
||||||
content: Some(ChatMessageContent::Text(res)),
|
|
||||||
reasoning_content: None,
|
|
||||||
refusal: None,
|
|
||||||
name: None,
|
|
||||||
audio: None,
|
|
||||||
tool_calls: None,
|
|
||||||
},
|
|
||||||
finish_reason: Some(FinishReason::StopSequenceReached),
|
|
||||||
logprobs: None,
|
|
||||||
}
|
|
||||||
};
|
|
||||||
response.choices.push(choice);
|
|
||||||
response
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn build_completion_response(
|
|
||||||
res: String,
|
|
||||||
model_name: &str,
|
|
||||||
completion_tokens: Option<u32>,
|
|
||||||
prompt_tokens: Option<u32>,
|
|
||||||
) -> ChatCompletionResponse {
|
|
||||||
let usage = if prompt_tokens.is_none() && completion_tokens.is_none() {
|
|
||||||
None
|
|
||||||
} else {
|
|
||||||
Some(Usage {
|
|
||||||
prompt_tokens,
|
|
||||||
prompt_secs: None,
|
|
||||||
completion_tokens,
|
|
||||||
completion_secs: None,
|
|
||||||
completion_per_token_secs: None,
|
|
||||||
completion_tps: None,
|
|
||||||
total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0),
|
|
||||||
prompt_tokens_details: None,
|
|
||||||
completion_tokens_details: None,
|
|
||||||
})
|
|
||||||
};
|
|
||||||
|
|
||||||
build_response(res, model_name, usage)
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn build_completion_response_with_time(
|
|
||||||
res: String,
|
|
||||||
model_name: &str,
|
|
||||||
completion_tokens: Option<u32>,
|
|
||||||
completion_secs: Option<f64>,
|
|
||||||
prompt_tokens: Option<u32>,
|
|
||||||
prompt_secs: Option<f64>,
|
|
||||||
) -> ChatCompletionResponse {
|
|
||||||
let usage = if prompt_tokens.is_none() && completion_tokens.is_none() {
|
|
||||||
None
|
|
||||||
} else {
|
|
||||||
let (completion_per_token_secs, completion_tps) = if let Some(completion_tokens) =
|
|
||||||
completion_tokens
|
|
||||||
&& let Some(completion_secs) = completion_secs
|
|
||||||
{
|
|
||||||
let per_token_secs = completion_secs / completion_tokens as f64;
|
|
||||||
let tps = completion_tokens as f64 / completion_secs;
|
|
||||||
(Some(per_token_secs), Some(tps))
|
|
||||||
} else {
|
|
||||||
(None, None)
|
|
||||||
};
|
|
||||||
Some(Usage {
|
|
||||||
prompt_tokens,
|
|
||||||
prompt_secs,
|
|
||||||
completion_tokens,
|
|
||||||
completion_secs,
|
|
||||||
completion_per_token_secs,
|
|
||||||
completion_tps,
|
|
||||||
total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0),
|
|
||||||
prompt_tokens_details: None,
|
|
||||||
completion_tokens_details: None,
|
|
||||||
})
|
|
||||||
};
|
|
||||||
|
|
||||||
build_response(res, model_name, usage)
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn build_completion_chunk_response(
|
|
||||||
res: String,
|
|
||||||
model_name: &str,
|
|
||||||
tool_call_id: Option<String>,
|
|
||||||
tool_call_content: Option<String>,
|
|
||||||
) -> ChatCompletionChunkResponse {
|
|
||||||
let id = uuid::Uuid::new_v4().to_string();
|
|
||||||
let mut response = ChatCompletionChunkResponse {
|
|
||||||
id: Some(id),
|
|
||||||
choices: vec![],
|
|
||||||
created: timestamp() as u32,
|
|
||||||
model: model_name.to_string(),
|
|
||||||
system_fingerprint: None,
|
|
||||||
object: "chat.completion.chunk".to_string(),
|
|
||||||
usage: None,
|
|
||||||
};
|
|
||||||
let choice = if let Some(tool_call_id) = tool_call_id {
|
|
||||||
let function = if let Some(content) = tool_call_content {
|
|
||||||
match serde_json::from_str::<serde_json::Value>(&content) {
|
|
||||||
Ok(json_value) => {
|
|
||||||
let name = json_value
|
|
||||||
.get("name")
|
|
||||||
.and_then(|v| v.as_str())
|
|
||||||
.map(|s| s.to_string());
|
|
||||||
|
|
||||||
let arguments = json_value.get("arguments").map(|v| v.to_string());
|
|
||||||
|
|
||||||
DeltaFunction { name, arguments }
|
|
||||||
}
|
|
||||||
Err(_) => DeltaFunction {
|
|
||||||
name: None,
|
|
||||||
arguments: Some(content),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
DeltaFunction {
|
|
||||||
name: None,
|
|
||||||
arguments: None,
|
|
||||||
}
|
|
||||||
};
|
|
||||||
ChatCompletionChunkChoice {
|
|
||||||
index: Some(0),
|
|
||||||
delta: DeltaChatMessage::Assistant {
|
|
||||||
content: None,
|
|
||||||
reasoning_content: None,
|
|
||||||
refusal: None,
|
|
||||||
name: None,
|
|
||||||
tool_calls: Some(vec![DeltaToolCall {
|
|
||||||
index: Some(0),
|
|
||||||
id: Some(tool_call_id),
|
|
||||||
r#type: Some("function".to_string()),
|
|
||||||
function,
|
|
||||||
}]),
|
|
||||||
},
|
|
||||||
finish_reason: None,
|
|
||||||
logprobs: None,
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
ChatCompletionChunkChoice {
|
|
||||||
index: Some(0),
|
|
||||||
delta: DeltaChatMessage::Assistant {
|
|
||||||
content: Some(ChatMessageContent::Text(res)),
|
|
||||||
reasoning_content: None,
|
|
||||||
refusal: None,
|
|
||||||
name: None,
|
|
||||||
tool_calls: None,
|
|
||||||
},
|
|
||||||
finish_reason: None,
|
|
||||||
logprobs: None,
|
|
||||||
}
|
|
||||||
};
|
|
||||||
response.choices.push(choice);
|
|
||||||
response
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn extract_mes(mes: &ChatCompletionParameters) -> Result<Vec<(String, String)>> {
|
pub fn extract_mes(mes: &ChatCompletionParameters) -> Result<Vec<(String, String)>> {
|
||||||
let mut mes_vec = Vec::new();
|
let mut mes_vec = Vec::new();
|
||||||
for chat_mes in mes.messages.clone() {
|
for chat_mes in mes.messages.clone() {
|
||||||
|
|||||||
@@ -0,0 +1,426 @@
|
|||||||
|
use crate::{
|
||||||
|
params::{
|
||||||
|
chat::{
|
||||||
|
AudioUrlType, ChatCompletionChoice, ChatCompletionChunkChoice,
|
||||||
|
ChatCompletionChunkResponse, ChatCompletionResponse, ChatMessage,
|
||||||
|
ChatMessageAudioContentPart, ChatMessageContent, ChatMessageContentPart,
|
||||||
|
ChatMessageImageContentPart, DeltaChatMessage, DeltaFunction, DeltaToolCall, Function,
|
||||||
|
ImageUrlType, ToolCall,
|
||||||
|
},
|
||||||
|
shared::{FinishReason, Usage},
|
||||||
|
},
|
||||||
|
utils::timestamp,
|
||||||
|
};
|
||||||
|
|
||||||
|
pub fn build_img_completion_response(
|
||||||
|
base64vec: &Vec<String>,
|
||||||
|
model_name: &str,
|
||||||
|
) -> ChatCompletionResponse {
|
||||||
|
let id = uuid::Uuid::new_v4().to_string();
|
||||||
|
let mut response = ChatCompletionResponse {
|
||||||
|
id: Some(id),
|
||||||
|
choices: vec![],
|
||||||
|
// created: chrono::Utc::now().timestamp() as u32,
|
||||||
|
created: timestamp() as u32,
|
||||||
|
model: model_name.to_string(),
|
||||||
|
service_tier: None,
|
||||||
|
system_fingerprint: None,
|
||||||
|
object: "chat.completion".to_string(),
|
||||||
|
usage: None,
|
||||||
|
};
|
||||||
|
let mut conten_part_vec = vec![];
|
||||||
|
for img_bas64 in base64vec {
|
||||||
|
let img_base64_prefix = "data:image/png;base64,".to_string() + img_bas64;
|
||||||
|
let part = ChatMessageContentPart::Image(ChatMessageImageContentPart {
|
||||||
|
r#type: "image".to_string(),
|
||||||
|
image_url: ImageUrlType {
|
||||||
|
url: img_base64_prefix,
|
||||||
|
detail: None,
|
||||||
|
},
|
||||||
|
});
|
||||||
|
conten_part_vec.push(part);
|
||||||
|
}
|
||||||
|
let choice = ChatCompletionChoice {
|
||||||
|
index: 0,
|
||||||
|
message: ChatMessage::Assistant {
|
||||||
|
content: Some(ChatMessageContent::ContentPart(conten_part_vec)),
|
||||||
|
reasoning_content: None,
|
||||||
|
refusal: None,
|
||||||
|
name: None,
|
||||||
|
audio: None,
|
||||||
|
tool_calls: None,
|
||||||
|
},
|
||||||
|
finish_reason: Some(FinishReason::StopSequenceReached),
|
||||||
|
logprobs: None,
|
||||||
|
};
|
||||||
|
response.choices.push(choice);
|
||||||
|
response
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn build_audio_completion_response(
|
||||||
|
base64_audio: &String,
|
||||||
|
model_name: &str,
|
||||||
|
) -> ChatCompletionResponse {
|
||||||
|
let id = uuid::Uuid::new_v4().to_string();
|
||||||
|
let mut response = ChatCompletionResponse {
|
||||||
|
id: Some(id),
|
||||||
|
choices: vec![],
|
||||||
|
created: timestamp() as u32,
|
||||||
|
model: model_name.to_string(),
|
||||||
|
service_tier: None,
|
||||||
|
system_fingerprint: None,
|
||||||
|
object: "chat.completion".to_string(),
|
||||||
|
usage: None,
|
||||||
|
};
|
||||||
|
|
||||||
|
let base64_audio = format!("data:audio/wav;base64,{}", base64_audio);
|
||||||
|
let conten_part_vec = vec![ChatMessageContentPart::Audio(ChatMessageAudioContentPart {
|
||||||
|
r#type: "audio".to_string(),
|
||||||
|
audio_url: AudioUrlType {
|
||||||
|
url: base64_audio.to_string(),
|
||||||
|
},
|
||||||
|
})];
|
||||||
|
let choice = ChatCompletionChoice {
|
||||||
|
index: 0,
|
||||||
|
message: ChatMessage::Assistant {
|
||||||
|
content: Some(ChatMessageContent::ContentPart(conten_part_vec)),
|
||||||
|
reasoning_content: None,
|
||||||
|
refusal: None,
|
||||||
|
name: None,
|
||||||
|
audio: None,
|
||||||
|
tool_calls: None,
|
||||||
|
},
|
||||||
|
finish_reason: Some(FinishReason::StopSequenceReached),
|
||||||
|
logprobs: None,
|
||||||
|
};
|
||||||
|
response.choices.push(choice);
|
||||||
|
response
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Builds a chat completion response from the model's output string.
|
||||||
|
///
|
||||||
|
/// This function handles two special formatting patterns in the input:
|
||||||
|
/// 1. Tool call formatting using the delimiter "<tool_call>" where the content before the first
|
||||||
|
/// delimiter is treated as the main content, and subsequent parts (separated by "</tool_call>")
|
||||||
|
/// are parsed as JSON tool call definitions.
|
||||||
|
/// 2. Reasoning content formatting where content between <think> and </think> tags is
|
||||||
|
/// extracted as reasoning_content, and the content after </think> becomes the main content.
|
||||||
|
///
|
||||||
|
/// # Arguments
|
||||||
|
/// * `res` - The raw response string from the model that may contain special formatting
|
||||||
|
/// * `model_name` - Name of the model generating the response
|
||||||
|
/// * `usage` - Optional usage statistics to include in the response
|
||||||
|
///
|
||||||
|
/// # Returns
|
||||||
|
/// A [ChatCompletionResponse](aha/src/params/chat.rs#L11-L32) with processed content, tool calls, and/or reasoning content
|
||||||
|
///
|
||||||
|
/// # Format specifications
|
||||||
|
/// - Tool call format: Content followed by "<tool_call>" and then JSON-formatted tool call data separated by "</tool_call>"
|
||||||
|
/// - Reasoning format: Content wrapped in <think> and </think> tags followed by actual response
|
||||||
|
fn build_response(res: String, model_name: &str, usage: Option<Usage>) -> ChatCompletionResponse {
|
||||||
|
let id = uuid::Uuid::new_v4().to_string();
|
||||||
|
let mut response = ChatCompletionResponse {
|
||||||
|
id: Some(id),
|
||||||
|
choices: vec![],
|
||||||
|
created: timestamp() as u32,
|
||||||
|
model: model_name.to_string(),
|
||||||
|
service_tier: None,
|
||||||
|
system_fingerprint: None,
|
||||||
|
object: "chat.completion".to_string(),
|
||||||
|
usage,
|
||||||
|
};
|
||||||
|
let (content, tool_calls) = if res.contains("<tool_call>") {
|
||||||
|
let mes: Vec<&str> = res.split("<tool_call>").collect();
|
||||||
|
let content = mes[0].to_string();
|
||||||
|
let mut tool_vec = Vec::new();
|
||||||
|
for (i, m) in mes.iter().enumerate().skip(1) {
|
||||||
|
let tool_mes = m.replace("</tool_call>", "");
|
||||||
|
let function = match serde_json::from_str::<serde_json::Value>(&tool_mes) {
|
||||||
|
Ok(json_value) => {
|
||||||
|
let name = json_value
|
||||||
|
.get("name")
|
||||||
|
.and_then(|v| v.as_str())
|
||||||
|
.map(|s| s.to_string())
|
||||||
|
.unwrap_or_default();
|
||||||
|
|
||||||
|
let arguments = json_value
|
||||||
|
.get("arguments")
|
||||||
|
.map(|v| v.to_string())
|
||||||
|
.unwrap_or_default();
|
||||||
|
|
||||||
|
Function { name, arguments }
|
||||||
|
}
|
||||||
|
Err(_) => Function {
|
||||||
|
name: "".to_string(),
|
||||||
|
arguments: "".to_string(),
|
||||||
|
},
|
||||||
|
};
|
||||||
|
let tool_call = ToolCall {
|
||||||
|
id: (i - 1).to_string(),
|
||||||
|
r#type: "function".to_string(),
|
||||||
|
function,
|
||||||
|
};
|
||||||
|
tool_vec.push(tool_call);
|
||||||
|
}
|
||||||
|
(content, Some(tool_vec))
|
||||||
|
} else {
|
||||||
|
(res, None)
|
||||||
|
};
|
||||||
|
let (content, reasoning_content) = if content.contains("</think>") {
|
||||||
|
let contents: Vec<&str> = content.split("</think>").collect();
|
||||||
|
let reasoning_content = contents[0].to_string().replace("<think>", "");
|
||||||
|
let content = contents[1].to_string();
|
||||||
|
(content, Some(reasoning_content))
|
||||||
|
} else {
|
||||||
|
(content, None)
|
||||||
|
};
|
||||||
|
let finish_reason = if tool_calls.is_some() {
|
||||||
|
Some(FinishReason::ToolCalls)
|
||||||
|
} else {
|
||||||
|
Some(FinishReason::StopSequenceReached)
|
||||||
|
};
|
||||||
|
let choice = ChatCompletionChoice {
|
||||||
|
index: 0,
|
||||||
|
message: ChatMessage::Assistant {
|
||||||
|
content: Some(ChatMessageContent::Text(content)),
|
||||||
|
reasoning_content,
|
||||||
|
refusal: None,
|
||||||
|
name: None,
|
||||||
|
audio: None,
|
||||||
|
tool_calls,
|
||||||
|
},
|
||||||
|
finish_reason,
|
||||||
|
logprobs: None,
|
||||||
|
};
|
||||||
|
response.choices.push(choice);
|
||||||
|
response
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn build_completion_response(
|
||||||
|
res: String,
|
||||||
|
model_name: &str,
|
||||||
|
completion_tokens: Option<u32>,
|
||||||
|
prompt_tokens: Option<u32>,
|
||||||
|
) -> ChatCompletionResponse {
|
||||||
|
let usage = if prompt_tokens.is_none() && completion_tokens.is_none() {
|
||||||
|
None
|
||||||
|
} else {
|
||||||
|
Some(Usage {
|
||||||
|
prompt_tokens,
|
||||||
|
prompt_secs: None,
|
||||||
|
completion_tokens,
|
||||||
|
completion_secs: None,
|
||||||
|
completion_per_token_secs: None,
|
||||||
|
completion_tps: None,
|
||||||
|
total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0),
|
||||||
|
prompt_tokens_details: None,
|
||||||
|
completion_tokens_details: None,
|
||||||
|
})
|
||||||
|
};
|
||||||
|
|
||||||
|
build_response(res, model_name, usage)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn build_completion_response_with_time(
|
||||||
|
res: String,
|
||||||
|
model_name: &str,
|
||||||
|
completion_tokens: Option<u32>,
|
||||||
|
completion_secs: Option<f64>,
|
||||||
|
prompt_tokens: Option<u32>,
|
||||||
|
prompt_secs: Option<f64>,
|
||||||
|
) -> ChatCompletionResponse {
|
||||||
|
let usage = if prompt_tokens.is_none() && completion_tokens.is_none() {
|
||||||
|
None
|
||||||
|
} else {
|
||||||
|
let (completion_per_token_secs, completion_tps) = if let Some(completion_tokens) =
|
||||||
|
completion_tokens
|
||||||
|
&& let Some(completion_secs) = completion_secs
|
||||||
|
{
|
||||||
|
let per_token_secs = completion_secs / completion_tokens as f64;
|
||||||
|
let tps = completion_tokens as f64 / completion_secs;
|
||||||
|
(Some(per_token_secs), Some(tps))
|
||||||
|
} else {
|
||||||
|
(None, None)
|
||||||
|
};
|
||||||
|
Some(Usage {
|
||||||
|
prompt_tokens,
|
||||||
|
prompt_secs,
|
||||||
|
completion_tokens,
|
||||||
|
completion_secs,
|
||||||
|
completion_per_token_secs,
|
||||||
|
completion_tps,
|
||||||
|
total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0),
|
||||||
|
prompt_tokens_details: None,
|
||||||
|
completion_tokens_details: None,
|
||||||
|
})
|
||||||
|
};
|
||||||
|
|
||||||
|
build_response(res, model_name, usage)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn build_chunk_response_with_usage(
|
||||||
|
model_name: &str,
|
||||||
|
completion_tokens: Option<u32>,
|
||||||
|
completion_secs: Option<f64>,
|
||||||
|
prompt_tokens: Option<u32>,
|
||||||
|
prompt_secs: Option<f64>,
|
||||||
|
) -> ChatCompletionChunkResponse {
|
||||||
|
let id = uuid::Uuid::new_v4().to_string();
|
||||||
|
let usage = if prompt_tokens.is_none() && completion_tokens.is_none() {
|
||||||
|
None
|
||||||
|
} else {
|
||||||
|
let (completion_per_token_secs, completion_tps) = if let Some(completion_tokens) =
|
||||||
|
completion_tokens
|
||||||
|
&& let Some(completion_secs) = completion_secs
|
||||||
|
{
|
||||||
|
let per_token_secs = completion_secs / completion_tokens as f64;
|
||||||
|
let tps = if (completion_secs - 0.0).abs() > 0.0001 {
|
||||||
|
completion_tokens as f64 / completion_secs
|
||||||
|
} else {
|
||||||
|
0.0
|
||||||
|
};
|
||||||
|
(Some(per_token_secs), Some(tps))
|
||||||
|
} else {
|
||||||
|
(None, None)
|
||||||
|
};
|
||||||
|
Some(Usage {
|
||||||
|
prompt_tokens,
|
||||||
|
prompt_secs,
|
||||||
|
completion_tokens,
|
||||||
|
completion_secs,
|
||||||
|
completion_per_token_secs,
|
||||||
|
completion_tps,
|
||||||
|
total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0),
|
||||||
|
prompt_tokens_details: None,
|
||||||
|
completion_tokens_details: None,
|
||||||
|
})
|
||||||
|
};
|
||||||
|
let mut response = ChatCompletionChunkResponse {
|
||||||
|
id: Some(id),
|
||||||
|
choices: vec![],
|
||||||
|
created: timestamp() as u32,
|
||||||
|
model: model_name.to_string(),
|
||||||
|
system_fingerprint: None,
|
||||||
|
object: "chat.completion.chunk".to_string(),
|
||||||
|
usage,
|
||||||
|
};
|
||||||
|
let choice = ChatCompletionChunkChoice {
|
||||||
|
index: Some(0),
|
||||||
|
delta: DeltaChatMessage::Assistant {
|
||||||
|
content: None,
|
||||||
|
reasoning_content: None,
|
||||||
|
refusal: None,
|
||||||
|
name: None,
|
||||||
|
tool_calls: None,
|
||||||
|
},
|
||||||
|
finish_reason: None,
|
||||||
|
logprobs: None,
|
||||||
|
};
|
||||||
|
response.choices.push(choice);
|
||||||
|
response
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn build_chunk_response_with_reasoning(
|
||||||
|
reasoning: String,
|
||||||
|
model_name: &str,
|
||||||
|
) -> ChatCompletionChunkResponse {
|
||||||
|
let id = uuid::Uuid::new_v4().to_string();
|
||||||
|
let mut response = ChatCompletionChunkResponse {
|
||||||
|
id: Some(id),
|
||||||
|
choices: vec![],
|
||||||
|
created: timestamp() as u32,
|
||||||
|
model: model_name.to_string(),
|
||||||
|
system_fingerprint: None,
|
||||||
|
object: "chat.completion.chunk".to_string(),
|
||||||
|
usage: None,
|
||||||
|
};
|
||||||
|
let choice = ChatCompletionChunkChoice {
|
||||||
|
index: Some(0),
|
||||||
|
delta: DeltaChatMessage::Assistant {
|
||||||
|
content: None,
|
||||||
|
reasoning_content: Some(reasoning),
|
||||||
|
refusal: None,
|
||||||
|
name: None,
|
||||||
|
tool_calls: None,
|
||||||
|
},
|
||||||
|
finish_reason: None,
|
||||||
|
logprobs: None,
|
||||||
|
};
|
||||||
|
response.choices.push(choice);
|
||||||
|
response
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn build_completion_chunk_response(
|
||||||
|
res: String,
|
||||||
|
model_name: &str,
|
||||||
|
tool_call_id: Option<String>,
|
||||||
|
tool_call_content: Option<String>,
|
||||||
|
) -> ChatCompletionChunkResponse {
|
||||||
|
let id = uuid::Uuid::new_v4().to_string();
|
||||||
|
let mut response = ChatCompletionChunkResponse {
|
||||||
|
id: Some(id),
|
||||||
|
choices: vec![],
|
||||||
|
created: timestamp() as u32,
|
||||||
|
model: model_name.to_string(),
|
||||||
|
system_fingerprint: None,
|
||||||
|
object: "chat.completion.chunk".to_string(),
|
||||||
|
usage: None,
|
||||||
|
};
|
||||||
|
let choice = if let Some(tool_call_id) = tool_call_id {
|
||||||
|
let function = if let Some(content) = tool_call_content {
|
||||||
|
match serde_json::from_str::<serde_json::Value>(&content) {
|
||||||
|
Ok(json_value) => {
|
||||||
|
let name = json_value
|
||||||
|
.get("name")
|
||||||
|
.and_then(|v| v.as_str())
|
||||||
|
.map(|s| s.to_string());
|
||||||
|
|
||||||
|
let arguments = json_value.get("arguments").map(|v| v.to_string());
|
||||||
|
|
||||||
|
DeltaFunction { name, arguments }
|
||||||
|
}
|
||||||
|
Err(_) => DeltaFunction {
|
||||||
|
name: None,
|
||||||
|
arguments: Some(content),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
DeltaFunction {
|
||||||
|
name: None,
|
||||||
|
arguments: None,
|
||||||
|
}
|
||||||
|
};
|
||||||
|
ChatCompletionChunkChoice {
|
||||||
|
index: Some(0),
|
||||||
|
delta: DeltaChatMessage::Assistant {
|
||||||
|
content: None,
|
||||||
|
reasoning_content: None,
|
||||||
|
refusal: None,
|
||||||
|
name: None,
|
||||||
|
tool_calls: Some(vec![DeltaToolCall {
|
||||||
|
index: Some(0),
|
||||||
|
id: Some(tool_call_id),
|
||||||
|
r#type: Some("function".to_string()),
|
||||||
|
function,
|
||||||
|
}]),
|
||||||
|
},
|
||||||
|
finish_reason: None,
|
||||||
|
logprobs: None,
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
ChatCompletionChunkChoice {
|
||||||
|
index: Some(0),
|
||||||
|
delta: DeltaChatMessage::Assistant {
|
||||||
|
content: Some(ChatMessageContent::Text(res)),
|
||||||
|
reasoning_content: None,
|
||||||
|
refusal: None,
|
||||||
|
name: None,
|
||||||
|
tool_calls: None,
|
||||||
|
},
|
||||||
|
finish_reason: None,
|
||||||
|
logprobs: None,
|
||||||
|
}
|
||||||
|
};
|
||||||
|
response.choices.push(choice);
|
||||||
|
response
|
||||||
|
}
|
||||||
@@ -89,17 +89,11 @@ fn deepseek_ocr_generate() -> Result<()> {
|
|||||||
let mut model = DeepseekOCRGenerateModel::init(&model_path, None, None)?;
|
let mut model = DeepseekOCRGenerateModel::init(&model_path, None, None)?;
|
||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
let i_start = Instant::now();
|
|
||||||
let res = model.generate(mes)?;
|
let res = model.generate(mes)?;
|
||||||
let i_duration = i_start.elapsed();
|
|
||||||
println!("generate: \n {:?}", res);
|
println!("generate: \n {:?}", res);
|
||||||
if let Some(usage) = &res.usage {
|
if let Some(usage) = &res.usage {
|
||||||
let num_token = usage.total_tokens;
|
println!("usage: \n {:?}", usage);
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -38,17 +38,11 @@ fn fun_asr_nano_generate() -> Result<()> {
|
|||||||
let mut fun_asr_model = FunAsrNanoGenerateModel::init(&model_path, None, None)?;
|
let mut fun_asr_model = FunAsrNanoGenerateModel::init(&model_path, None, None)?;
|
||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
let i_start = Instant::now();
|
|
||||||
let res = fun_asr_model.generate(mes)?;
|
let res = fun_asr_model.generate(mes)?;
|
||||||
let i_duration = i_start.elapsed();
|
|
||||||
println!("generate: \n {:?}", res);
|
println!("generate: \n {:?}", res);
|
||||||
if let Some(usage) = &res.usage {
|
if let Some(usage) = &res.usage {
|
||||||
let num_token = usage.total_tokens;
|
println!("usage: \n {:?}", usage);
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -62,16 +62,10 @@ fn gelab_zero_generate() -> Result<()> {
|
|||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
|
||||||
let i_start = Instant::now();
|
|
||||||
let res = qwen3vl.generate(mes)?;
|
let res = qwen3vl.generate(mes)?;
|
||||||
let i_duration = i_start.elapsed();
|
|
||||||
println!("generate: \n {:?}", res);
|
println!("generate: \n {:?}", res);
|
||||||
if let Some(usage) = &res.usage {
|
if let Some(usage) = &res.usage {
|
||||||
let num_token = usage.total_tokens;
|
println!("usage: \n {:?}", usage);
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -80,16 +80,10 @@ fn gguf_test() -> Result<()> {
|
|||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
|
||||||
let i_start = Instant::now();
|
|
||||||
let res = gguf_qwen3_5.generate(mes)?;
|
let res = gguf_qwen3_5.generate(mes)?;
|
||||||
let i_duration = i_start.elapsed();
|
|
||||||
println!("generate: \n {:?}", res);
|
println!("generate: \n {:?}", res);
|
||||||
if let Some(usage) = &res.usage {
|
if let Some(usage) = &res.usage {
|
||||||
let num_token = usage.total_tokens;
|
println!("usage: \n {:?}", usage);
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -39,23 +39,17 @@ fn glm_asr_nano_generate() -> Result<()> {
|
|||||||
let mut glm_asr_model = GlmAsrNanoGenerateModel::init(&model_path, None, None)?;
|
let mut glm_asr_model = GlmAsrNanoGenerateModel::init(&model_path, None, None)?;
|
||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
let i_start = Instant::now();
|
|
||||||
let res = glm_asr_model.generate(mes)?;
|
let res = glm_asr_model.generate(mes)?;
|
||||||
let i_duration = i_start.elapsed();
|
|
||||||
println!("generate: \n {:?}", res);
|
println!("generate: \n {:?}", res);
|
||||||
if let Some(usage) = &res.usage {
|
if let Some(usage) = &res.usage {
|
||||||
let num_token = usage.total_tokens;
|
println!("usage: \n {:?}", usage);
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn glm_asr_nano_stream() -> Result<()> {
|
async fn glm_asr_nano_stream() -> Result<()> {
|
||||||
// RUST_BACKTRACE=1 cargo test -F cuda glm_asr_nano_stream -r -- --nocapture
|
// RUST_BACKTRACE=1 cargo test -F cuda --test test_glm_asr_nano glm_asr_nano_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"))?;
|
||||||
let model_path = format!("{}/ZhipuAI/GLM-ASR-Nano-2512/", save_dir);
|
let model_path = format!("{}/ZhipuAI/GLM-ASR-Nano-2512/", save_dir);
|
||||||
|
|||||||
@@ -39,17 +39,11 @@ fn glm_ocr_generate() -> Result<()> {
|
|||||||
let mut model = GlmOcrGenerateModel::init(&model_path, None, None)?;
|
let mut model = GlmOcrGenerateModel::init(&model_path, None, None)?;
|
||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
let i_start = Instant::now();
|
|
||||||
let res = model.generate(mes)?;
|
let res = model.generate(mes)?;
|
||||||
let i_duration = i_start.elapsed();
|
|
||||||
println!("generate: \n {:?}", res);
|
println!("generate: \n {:?}", res);
|
||||||
if let Some(usage) = &res.usage {
|
if let Some(usage) = &res.usage {
|
||||||
let num_token = usage.total_tokens;
|
println!("usage: \n {:?}", usage);
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,76 +0,0 @@
|
|||||||
use aha::models::common::model_mapping::WhichModel;
|
|
||||||
|
|
||||||
// Import helper functions from api module - these will need to be made public
|
|
||||||
// or tested through integration testing
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_model_type_classification() {
|
|
||||||
// Since get_model_type and get_model_id are private to api.rs,
|
|
||||||
// we document the expected behavior here for reference:
|
|
||||||
//
|
|
||||||
// LLM models: MiniCPM4_0_5B, Qwen2_5VL3B, Qwen2_5VL7B, Qwen3_0_6B,
|
|
||||||
// Qwen3VL2B, Qwen3VL4B, Qwen3VL8B, Qwen3VL32B
|
|
||||||
// OCR models: DeepSeekOCR, HunyuanOCR, PaddleOCRVL
|
|
||||||
// ASR models: Qwen3ASR0_6B, Qwen3ASR1_7B, GlmASRNano2512, FunASRNano2512
|
|
||||||
// Image models: RMBG2_0, VoxCPM, VoxCPM1_5
|
|
||||||
|
|
||||||
// This test documents the expected model type classification
|
|
||||||
let llm_models = [
|
|
||||||
WhichModel::MiniCPM4_0_5B,
|
|
||||||
WhichModel::Qwen2_5VL3B,
|
|
||||||
WhichModel::Qwen2_5VL7B,
|
|
||||||
WhichModel::Qwen3_0_6B,
|
|
||||||
WhichModel::Qwen3VL2B,
|
|
||||||
WhichModel::Qwen3VL4B,
|
|
||||||
WhichModel::Qwen3VL8B,
|
|
||||||
WhichModel::Qwen3VL32B,
|
|
||||||
];
|
|
||||||
|
|
||||||
let ocr_models = [
|
|
||||||
WhichModel::DeepSeekOCR,
|
|
||||||
WhichModel::HunyuanOCR,
|
|
||||||
WhichModel::PaddleOCRVL,
|
|
||||||
];
|
|
||||||
|
|
||||||
let asr_models = [
|
|
||||||
WhichModel::Qwen3ASR0_6B,
|
|
||||||
WhichModel::Qwen3ASR1_7B,
|
|
||||||
WhichModel::GlmASRNano2512,
|
|
||||||
WhichModel::FunASRNano2512,
|
|
||||||
];
|
|
||||||
|
|
||||||
let image_models = [
|
|
||||||
WhichModel::RMBG2_0,
|
|
||||||
WhichModel::VoxCPM,
|
|
||||||
WhichModel::VoxCPM1_5,
|
|
||||||
];
|
|
||||||
|
|
||||||
// Verify counts
|
|
||||||
assert_eq!(llm_models.len(), 8);
|
|
||||||
assert_eq!(ocr_models.len(), 3);
|
|
||||||
assert_eq!(asr_models.len(), 4);
|
|
||||||
assert_eq!(image_models.len(), 3);
|
|
||||||
|
|
||||||
// Total models
|
|
||||||
assert_eq!(
|
|
||||||
llm_models.len() + ocr_models.len() + asr_models.len() + image_models.len(),
|
|
||||||
18
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
// Note: Integration tests for the /health and /models endpoints
|
|
||||||
// should be done with a running server. These would typically:
|
|
||||||
//
|
|
||||||
// 1. Start the server with a test model
|
|
||||||
// 2. Make HTTP requests to /health and /models
|
|
||||||
// 3. Verify the response format and status codes
|
|
||||||
//
|
|
||||||
// Example (pseudo-code):
|
|
||||||
//
|
|
||||||
// #[tokio::test]
|
|
||||||
// async fn test_health_endpoint() {
|
|
||||||
// let resp = reqwest::get("http://localhost:10100/health").await.unwrap();
|
|
||||||
// assert_eq!(resp.status(), 200);
|
|
||||||
// let json: serde_json::Value = resp.json().await.unwrap();
|
|
||||||
// assert_eq!(json["status"], "ok");
|
|
||||||
// }
|
|
||||||
@@ -39,18 +39,11 @@ fn hunyuan_ocr_generate() -> Result<()> {
|
|||||||
let mut model = HunyuanOCRGenerateModel::init(&model_path, None, None)?;
|
let mut model = HunyuanOCRGenerateModel::init(&model_path, None, None)?;
|
||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
let i_start = Instant::now();
|
|
||||||
let res = model.generate(mes)?;
|
let res = model.generate(mes)?;
|
||||||
let i_duration = i_start.elapsed();
|
|
||||||
println!("generate: \n {:?}", res);
|
println!("generate: \n {:?}", res);
|
||||||
if let Some(usage) = &res.usage {
|
if let Some(usage) = &res.usage {
|
||||||
let num_token = usage.total_tokens;
|
println!("usage: \n {:?}", usage);
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+4
-10
@@ -31,17 +31,11 @@ fn lfm2_generate() -> Result<()> {
|
|||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
|
||||||
let i_start = Instant::now();
|
let res = model.generate(mes)?;
|
||||||
let result = model.generate(mes)?;
|
println!("generate: \n {:?}", res);
|
||||||
let i_duration = i_start.elapsed();
|
if let Some(usage) = &res.usage {
|
||||||
println!("generate: \n {:?}", result);
|
println!("usage: \n {:?}", usage);
|
||||||
if let Some(usage) = &result.usage {
|
|
||||||
let num_token = usage.total_tokens;
|
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
+4
-11
@@ -43,18 +43,11 @@ fn lfm2vl_generate() -> Result<()> {
|
|||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
|
||||||
let i_start = Instant::now();
|
let res = model.generate(mes)?;
|
||||||
let result = model.generate(mes)?;
|
println!("generate: \n {:?}", res);
|
||||||
let i_duration = i_start.elapsed();
|
if let Some(usage) = &res.usage {
|
||||||
println!("generate: \n {:?}", result);
|
println!("usage: \n {:?}", usage);
|
||||||
if let Some(usage) = &result.usage {
|
|
||||||
let num_token = usage.total_tokens;
|
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+4
-11
@@ -33,18 +33,11 @@ fn minicpm_generate() -> Result<()> {
|
|||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
|
||||||
let i_start = Instant::now();
|
let res = model.generate(mes)?;
|
||||||
let result = model.generate(mes)?;
|
println!("generate: \n {:?}", res);
|
||||||
let i_duration = i_start.elapsed();
|
if let Some(usage) = &res.usage {
|
||||||
println!("generate: \n {:?}", result);
|
println!("usage: \n {:?}", usage);
|
||||||
if let Some(usage) = &result.usage {
|
|
||||||
let num_token = usage.total_tokens;
|
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -89,17 +89,11 @@ fn paddleocr_vl_generate() -> Result<()> {
|
|||||||
let mut model = PaddleOCRVLGenerateModel::init(&model_path, None, None)?;
|
let mut model = PaddleOCRVLGenerateModel::init(&model_path, None, None)?;
|
||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
let i_start = Instant::now();
|
|
||||||
let res = model.generate(mes)?;
|
let res = model.generate(mes)?;
|
||||||
let i_duration = i_start.elapsed();
|
|
||||||
println!("generate: \n {:?}", res);
|
println!("generate: \n {:?}", res);
|
||||||
if let Some(usage) = &res.usage {
|
if let Some(usage) = &res.usage {
|
||||||
let num_token = usage.total_tokens;
|
println!("usage: \n {:?}", usage);
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+4
-11
@@ -46,18 +46,11 @@ fn qwen2_5vl_generate() -> Result<()> {
|
|||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
|
||||||
let i_start = Instant::now();
|
let res = model.generate(mes)?;
|
||||||
let result = model.generate(mes)?;
|
println!("generate: \n {:?}", res);
|
||||||
let i_duration = i_start.elapsed();
|
if let Some(usage) = &res.usage {
|
||||||
println!("generate: \n {:?}", result);
|
println!("usage: \n {:?}", usage);
|
||||||
if let Some(usage) = &result.usage {
|
|
||||||
let num_token = usage.total_tokens;
|
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+8
-12
@@ -21,7 +21,8 @@ fn qwen3_0_6b_generate() -> Result<()> {
|
|||||||
"role": "user",
|
"role": "user",
|
||||||
"content": "你好啊,你是谁"
|
"content": "你好啊,你是谁"
|
||||||
}
|
}
|
||||||
]
|
],
|
||||||
|
"enable_thinking": true
|
||||||
}
|
}
|
||||||
"#;
|
"#;
|
||||||
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
||||||
@@ -30,17 +31,11 @@ fn qwen3_0_6b_generate() -> Result<()> {
|
|||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
|
||||||
let i_start = Instant::now();
|
let res = model.generate(mes)?;
|
||||||
let result = model.generate(mes)?;
|
println!("generate: \n {:?}", res);
|
||||||
let i_duration = i_start.elapsed();
|
if let Some(usage) = &res.usage {
|
||||||
println!("generate: \n {:?}", result);
|
println!("usage: \n {:?}", usage);
|
||||||
if let Some(usage) = &result.usage {
|
|
||||||
let num_token = usage.total_tokens;
|
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
@@ -61,7 +56,8 @@ async fn qwen3_0_6b_stream() -> Result<()> {
|
|||||||
"role": "user",
|
"role": "user",
|
||||||
"content": "你是谁"
|
"content": "你是谁"
|
||||||
}
|
}
|
||||||
]
|
],
|
||||||
|
"enable_thinking": true
|
||||||
}
|
}
|
||||||
"#;
|
"#;
|
||||||
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
||||||
|
|||||||
+8
-11
@@ -33,26 +33,22 @@ fn qwen3_5_generate() -> Result<()> {
|
|||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
]
|
],
|
||||||
|
"enable_thinking": true
|
||||||
}
|
}
|
||||||
"#;
|
"#;
|
||||||
|
// "metadata": {"enable_thinking": "true"}
|
||||||
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
||||||
let i_start = Instant::now();
|
let i_start = Instant::now();
|
||||||
let mut qwen3_5 = Qwen3_5GenerateModel::init(&model_path, None, None)?;
|
let mut qwen3_5 = Qwen3_5GenerateModel::init(&model_path, None, None)?;
|
||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
|
||||||
let i_start = Instant::now();
|
|
||||||
let res = qwen3_5.generate(mes)?;
|
let res = qwen3_5.generate(mes)?;
|
||||||
let i_duration = i_start.elapsed();
|
|
||||||
println!("generate: \n {:?}", res);
|
println!("generate: \n {:?}", res);
|
||||||
if let Some(usage) = &res.usage {
|
if let Some(usage) = &res.usage {
|
||||||
let num_token = usage.total_tokens;
|
println!("usage: \n {:?}", usage);
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -75,16 +71,17 @@ async fn qwen3_5_stream() -> Result<()> {
|
|||||||
"type": "image",
|
"type": "image",
|
||||||
"image_url":
|
"image_url":
|
||||||
{
|
{
|
||||||
"url": "file:///home/jhq/Downloads/gougou1.jpg"
|
"url": "file://./assets/img/ocr_test3.png"
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"type": "text",
|
"type": "text",
|
||||||
"text": "描述这张图片."
|
"text": "OCR"
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
]
|
],
|
||||||
|
"enable_thinking": true
|
||||||
}
|
}
|
||||||
"#;
|
"#;
|
||||||
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
||||||
|
|||||||
@@ -35,17 +35,11 @@ fn qwen3_asr_generate() -> Result<()> {
|
|||||||
let mut model = Qwen3AsrGenerateModel::init(&model_path, None, None)?;
|
let mut model = Qwen3AsrGenerateModel::init(&model_path, None, None)?;
|
||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
let i_start = Instant::now();
|
|
||||||
let res = model.generate(mes)?;
|
let res = model.generate(mes)?;
|
||||||
let i_duration = i_start.elapsed();
|
|
||||||
println!("generate: \n {:?}", res);
|
println!("generate: \n {:?}", res);
|
||||||
if let Some(usage) = &res.usage {
|
if let Some(usage) = &res.usage {
|
||||||
let num_token = usage.total_tokens;
|
println!("usage: \n {:?}", usage);
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -94,17 +94,11 @@ fn qwen3vl_generate() -> Result<()> {
|
|||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
|
||||||
let i_start = Instant::now();
|
|
||||||
let res = qwen3vl.generate(mes)?;
|
let res = qwen3vl.generate(mes)?;
|
||||||
let i_duration = i_start.elapsed();
|
|
||||||
println!("generate: \n {:?}", res);
|
println!("generate: \n {:?}", res);
|
||||||
if let Some(usage) = &res.usage {
|
if let Some(usage) = &res.usage {
|
||||||
let num_token = usage.total_tokens;
|
println!("usage: \n {:?}", usage);
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -34,17 +34,10 @@ fn robo_brain_generate() -> Result<()> {
|
|||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||||
|
|
||||||
let i_start = Instant::now();
|
let res = model.generate(mes)?;
|
||||||
let result = model.generate(mes)?;
|
println!("generate: \n {:?}", res);
|
||||||
let i_duration = i_start.elapsed();
|
if let Some(usage) = &res.usage {
|
||||||
println!("generate: \n {:?}", result);
|
println!("usage: \n {:?}", usage);
|
||||||
if let Some(usage) = &result.usage {
|
|
||||||
let num_token = usage.total_tokens;
|
|
||||||
let duration_secs = i_duration.as_secs_f64();
|
|
||||||
let tps = num_token as f64 / duration_secs;
|
|
||||||
println!("Tokens per second (TPS): {:.2}", tps);
|
|
||||||
}
|
}
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user