From bd8eee65207772638d3845b652770c700b1894d8 Mon Sep 17 00:00:00 2001
From: jhqxxx <18280426169@163.com>
Date: Thu, 2 Apr 2026 22:22:52 +0800
Subject: [PATCH] refactor generate code
---
README.md | 14 +-
README.zh-CN.md | 12 +-
docs/changelog.md | 5 +
docs/changelog.zh-CN.md | 5 +
src/chat_template/mod.rs | 3 +
src/models/common/generate.rs | 86 +++++-
src/models/common/mod.rs | 6 +-
src/models/deepseek_ocr/generate.rs | 1 +
src/models/fun_asr_nano/generate.rs | 54 +---
src/models/fun_asr_nano/model.rs | 2 +-
src/models/fun_asr_nano/processor.rs | 5 +-
src/models/glm_asr_nano/generate.rs | 146 +++------
src/models/glm_asr_nano/model.rs | 51 +++-
src/models/glm_asr_nano/processor.rs | 3 +
src/models/glm_ocr/generate.rs | 171 ++++-------
src/models/glm_ocr/model.rs | 46 ++-
src/models/hunyuan_ocr/config.rs | 2 +-
src/models/hunyuan_ocr/generate.rs | 169 ++++-------
src/models/hunyuan_ocr/model.rs | 56 +++-
src/models/lfm2/generate.rs | 119 +++-----
src/models/lfm2/model.rs | 28 +-
src/models/lfm2vl/generate.rs | 166 ++++-------
src/models/lfm2vl/model.rs | 54 +++-
src/models/minicpm4/generate.rs | 126 +++-----
src/models/minicpm4/model.rs | 35 ++-
src/models/paddleocr_vl/generate.rs | 163 ++++------
src/models/paddleocr_vl/model.rs | 72 ++++-
src/models/qwen2_5vl/generate.rs | 42 ++-
src/models/qwen3/generate.rs | 138 +++------
src/models/qwen3/model.rs | 23 +-
src/models/qwen3_5/generate.rs | 250 ++++------------
src/models/qwen3_5/model.rs | 47 ++-
src/models/qwen3_asr/config.rs | 2 +-
src/models/qwen3_asr/generate.rs | 39 ++-
src/models/qwen3vl/generate.rs | 331 ++++++++++-----------
src/models/qwen3vl/model.rs | 56 +++-
src/models/rmbg2_0/generate.rs | 3 +-
src/models/voxcpm/generate.rs | 4 +-
src/params/chat.rs | 2 +
src/utils/mod.rs | 310 +------------------
src/utils/response_utils.rs | 426 +++++++++++++++++++++++++++
tests/test_deepseek_ocr.rs | 8 +-
tests/test_fun_asr_nano.rs | 8 +-
tests/test_gelab_zero.rs | 8 +-
tests/test_gguf_qwen3_5.rs | 8 +-
tests/test_glm_asr_nano.rs | 10 +-
tests/test_glm_ocr.rs | 8 +-
tests/test_health_models.rs | 76 -----
tests/test_hunyuan_ocr.rs | 9 +-
tests/test_lfm2.rs | 14 +-
tests/test_lfm2vl.rs | 15 +-
tests/test_minicpm4.rs | 15 +-
tests/test_paddleocr_vl.rs | 8 +-
tests/test_qwen2_5vl.rs | 15 +-
tests/test_qwen3.rs | 20 +-
tests/test_qwen3_5.rs | 19 +-
tests/test_qwen3_asr.rs | 8 +-
tests/test_qwen3vl.rs | 8 +-
tests/test_robo_brain.rs | 15 +-
59 files changed, 1745 insertions(+), 1800 deletions(-)
create mode 100644 src/utils/response_utils.rs
delete mode 100644 tests/test_health_models.rs
diff --git a/README.md b/README.md
index e402dcc..dc9599a 100644
--- a/README.md
+++ b/README.md
@@ -29,11 +29,11 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
| Category | Models |
|----------|--------|
-| **Text** | Qwen3, MiniCPM4,
LFM2, LFM2.5 |
+| **Text** | Qwen3, MiniCPM4, LFM2, LFM2.5 |
| **Vision** | Qwen2.5-VL, Qwen3-VL, Qwen3.5,
LFM2.5-VL, LFM2-VL |
-| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2 ,
PaddleOCR-VL, PaddleOCR-VL1.5,
Hunyuan-OCR, GLM-OCR |
+| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2 , PaddleOCR-VL
PaddleOCR-VL1.5, Hunyuan-OCR, GLM-OCR |
| **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR |
-| **Audio** | VoxCPM, VoxCPM1.5 |
+| **TTS** | VoxCPM, VoxCPM1.5 |
| **Image** | RMBG-2.0 (background removal) |
## 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
## Changelog
+### 2026-04-02
+- refactor generate code
+- \...\ The content of the thought chain is returned using the reasoning_content field.
+- chat response add time info
+
### 2026-04-01
- 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-1.2B
-### v0.2.3 (2026-03-18)
-- add DeepSeek-OCR-2
-
**[View full changelog](docs/changelog.md)** →
diff --git a/README.zh-CN.md b/README.zh-CN.md
index a624d72..386e946 100644
--- a/README.zh-CN.md
+++ b/README.zh-CN.md
@@ -28,11 +28,11 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
| 类别 | 模型 |
|------|------|
-| **文本** | Qwen3, MiniCPM4,
LFM2, LFM2.5 |
+| **文本** | Qwen3, MiniCPM4, LFM2, LFM2.5 |
| **视觉** | Qwen2.5-VL, Qwen3-VL, Qwen3.5
LFM2.5-VL, LFM2-VL |
-| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2 ,
PaddleOCR-VL, PaddleOCR-VL1.5,
Hunyuan-OCR, GLM-OCR |
+| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2, PaddleOCR-VL,
PaddleOCR-VL1.5, Hunyuan-OCR, GLM-OCR |
| **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR |
-| **音频** | VoxCPM, VoxCPM1.5 |
+| **TTS** | VoxCPM, VoxCPM1.5 |
| **图像** | RMBG-2.0 (背景移除) |
## 为什么选择 aha?
@@ -45,6 +45,12 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
- **🧠 注意力优化** - 可选 Flash Attention 支持,优化长序列处理
## 更新日志
+## Changelog
+### 2026-04-02
+- 重构生成代码
+- \...\ 思维链内容使用reasoning_content字段返回。
+- 对话返回添加耗时信息
+
### 2026-04-01
- 重构 deepseek_ocr/fun_asr_nano 生成代码
diff --git a/docs/changelog.md b/docs/changelog.md
index 5a3bc11..bc980ed 100644
--- a/docs/changelog.md
+++ b/docs/changelog.md
@@ -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/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
+### 2026-04-02
+- refactor generate code
+- \...\ The content of the thought chain is returned using the reasoning_content field.
+- response add time info
+
### 2026-04-01
- refactor deepseek_ocr/fun_asr_nano generate code
diff --git a/docs/changelog.zh-CN.md b/docs/changelog.zh-CN.md
index 8585629..4d5bf0c 100644
--- a/docs/changelog.zh-CN.md
+++ b/docs/changelog.zh-CN.md
@@ -5,6 +5,11 @@
格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/),
本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。
+### 2026-04-02
+- 重构生成代码
+- \...\ 思维链内容使用reasoning_content字段返回。
+- 对话返回添加耗时信息
+
### 2026-04-01
- 重构 deepseek_ocr/fun_asr_nano 生成代码
diff --git a/src/chat_template/mod.rs b/src/chat_template/mod.rs
index 5a3b269..c680d5f 100644
--- a/src/chat_template/mod.rs
+++ b/src/chat_template/mod.rs
@@ -133,6 +133,9 @@ impl<'a> ChatTemplate<'a> {
pub fn apply_chat_template(&self, messages: &ChatCompletionParameters) -> Result {
let enable_thinking = extract_metadata_value::(&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! {
messages => &messages.messages,
tools => &messages.tools.as_ref(),
diff --git a/src/models/common/generate.rs b/src/models/common/generate.rs
index eb9c079..1f34d8b 100644
--- a/src/models/common/generate.rs
+++ b/src/models/common/generate.rs
@@ -9,7 +9,10 @@ use crate::{
models::common::{InferenceModel, MultiModalData},
params::chat::{ChatCompletionChunkResponse, ChatCompletionResponse},
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(
temperature: Option,
@@ -98,6 +101,17 @@ fn sample_and_push(
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(
model: &mut M,
tokenizer: &TokenizerModel,
@@ -155,6 +169,7 @@ pub fn generate_stream_generic(
top_k: Option,
seed: u64,
max_tokens: u32,
+ in_reasoning: bool,
device: &Device,
model_name: &str,
) -> Result>> {
@@ -167,14 +182,23 @@ pub fn generate_stream_generic(
max_tokens,
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 eos_ids = model.stop_token_ids();
let stream = stream! {
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 错误累积
for _ in 0..ctx.sample_len {
+ let i_start = Instant::now();
let logits = if ctx.seqlen_offset == 0 {
model.forward_initial(&input_ids, ctx.seqlen_offset, data.clone())
+
} else {
model.forward_step(&input_ids, ctx.seqlen_offset)
}?;
@@ -183,6 +207,13 @@ pub fn generate_stream_generic(
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
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() {
@@ -204,9 +235,60 @@ pub fn generate_stream_generic(
continue;
}
error_tokens.clear();
- yield Ok(build_completion_chunk_response(decoded, model_name, None, None));
+ if decoded.eq("") {
+ in_reasoning = true;
+ input_ids = ctx.prepare_for_next_token(next_token)?;
+ continue;
+ }
+ if decoded.eq("") {
+ in_reasoning = false;
+ input_ids = ctx.prepare_for_next_token(next_token)?;
+ continue;
+ }
+ // 处理特殊标记和工具调用
+ match decoded.as_str() {
+ "" => {
+ // 开始工具调用
+ tool_call_id = Some(uuid::Uuid::new_v4().to_string());
+ input_ids = ctx.prepare_for_next_token(next_token)?;
+ continue;
+ }
+ "" => {
+ // 结束工具调用
+ 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) {
+ yield Ok(build_chunk_response_with_usage(model_name, completion_tokens.into(), completion_secs.into(), prompt_tokens.into(), prompt_secs.into()));
break;
}
input_ids = ctx.prepare_for_next_token(next_token)?;
diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs
index 7065ec7..19c19a0 100644
--- a/src/models/common/mod.rs
+++ b/src/models/common/mod.rs
@@ -18,14 +18,18 @@ impl MultiModalData {
}
}
+#[allow(unused)]
pub trait InferenceModel {
/// 初始前向传播(考虑多模态输入)
+ /// 默认实现无特殊数据
fn forward_initial(
&mut self,
input_ids: &Tensor,
seqlen_offset: usize,
data: MultiModalData,
- ) -> Result;
+ ) -> Result {
+ Self::forward_step(self, input_ids, seqlen_offset)
+ }
/// 后续前向传播(自回归步骤)
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result;
diff --git a/src/models/deepseek_ocr/generate.rs b/src/models/deepseek_ocr/generate.rs
index 79eee51..939afaf 100644
--- a/src/models/deepseek_ocr/generate.rs
+++ b/src/models/deepseek_ocr/generate.rs
@@ -171,6 +171,7 @@ impl GenerateModel for DeepseekOCRGenerateModel {
None,
seed,
max_tokens,
+ false,
&self.device,
&self.model_name,
)?;
diff --git a/src/models/fun_asr_nano/generate.rs b/src/models/fun_asr_nano/generate.rs
index f97c04c..446648d 100644
--- a/src/models/fun_asr_nano/generate.rs
+++ b/src/models/fun_asr_nano/generate.rs
@@ -30,8 +30,6 @@ pub struct FunAsrNanoGenerateModel {
fun_asr_nano: FunAsrNanoModel,
device: Device,
dtype: DType,
- // eos_token_id1: u32,
- // eos_token_id2: u32,
generation_config: Qwen3GenerationConfig,
model_name: String,
}
@@ -90,8 +88,6 @@ impl FunAsrNanoGenerateModel {
fun_asr_nano,
device,
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,
model_name,
})
@@ -165,58 +161,10 @@ impl GenerateModel for FunAsrNanoGenerateModel {
top_k.into(),
seed,
max_tokens,
+ false,
&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 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)))
}
}
diff --git a/src/models/fun_asr_nano/model.rs b/src/models/fun_asr_nano/model.rs
index 2d3118e..484e4b2 100644
--- a/src/models/fun_asr_nano/model.rs
+++ b/src/models/fun_asr_nano/model.rs
@@ -611,7 +611,7 @@ impl FunAsrNanoModel {
config.audio_adaptor_conf.n_layer,
8,
)?;
- let llm = Qwen3Model::new(llm_cfg, vb.pp("llm"))?;
+ let llm = Qwen3Model::new(llm_cfg, vb.pp("llm"), vec![])?;
Ok(Self {
audio_encoder,
audio_adaptor,
diff --git a/src/models/fun_asr_nano/processor.rs b/src/models/fun_asr_nano/processor.rs
index 662409b..04adcff 100644
--- a/src/models/fun_asr_nano/processor.rs
+++ b/src/models/fun_asr_nano/processor.rs
@@ -1,5 +1,5 @@
use crate::params::chat::ChatCompletionParameters;
-use anyhow::Result;
+use anyhow::{Result, anyhow};
use candle_core::{D, Device, Tensor};
use crate::{
@@ -92,6 +92,9 @@ impl FunAsrNanoProcessor {
source_ids.extend_from_slice(&sub_token);
fbank_mask.extend_from_slice(&vec![0u32; sub_token.len()]);
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 (speech, speech_lengths) = self.extract_fbank(audio)?;
let olens = 1 + (speech_lengths - 3 + 2) / 2;
diff --git a/src/models/glm_asr_nano/generate.rs b/src/models/glm_asr_nano/generate.rs
index a28987e..66eada0 100644
--- a/src/models/glm_asr_nano/generate.rs
+++ b/src/models/glm_asr_nano/generate.rs
@@ -1,11 +1,13 @@
use crate::{
- models::common::generate::get_logit_processor,
+ models::common::{
+ MultiModalData,
+ generate::{GenerationContext, generate_generic, generate_stream_generic},
+ },
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
};
-use anyhow::{Result, anyhow};
+use anyhow::Result;
use candle_core::{DType, Device, Tensor};
use candle_nn::VarBuilder;
-use rocket::async_stream::stream;
use rocket::futures::Stream;
use crate::{
@@ -17,10 +19,7 @@ use crate::{
},
},
tokenizer::TokenizerModel,
- utils::{
- build_completion_chunk_response, build_completion_response, find_type_files, get_device,
- get_dtype,
- },
+ utils::{find_type_files, get_device, get_dtype},
};
pub struct GlmAsrNanoGenerateModel<'a> {
@@ -30,9 +29,6 @@ pub struct GlmAsrNanoGenerateModel<'a> {
glm_asr_nano: GlmAsrNanoModel,
device: Device,
dtype: DType,
- eos_token_id1: u32,
- eos_token_id2: u32,
- eos_token_id3: u32,
model_name: String,
}
@@ -48,7 +44,8 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
let dtype = get_dtype(dtype, cfg_dtype);
let model_list = find_type_files(path, "safetensors")?;
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)
.file_name()
.and_then(|s| s.to_str())
@@ -61,9 +58,6 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
glm_asr_nano,
device,
dtype,
- eos_token_id1: 59246,
- eos_token_id2: 59253,
- eos_token_id3: 59255,
model_name,
})
}
@@ -72,46 +66,33 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result {
let seed = mes.seed.unwrap_or(34562) as u64;
- let 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 (input_features, audio_token_lengths, replace_text) =
self.processor.process_info(&mes, &render_text)?;
- let mut input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
- let mut input_features = Some(input_features.to_dtype(self.dtype)?);
- let mut audio_token_lengths = Some(audio_token_lengths);
- let mut seq_len = input_ids.dim(1)?;
- let prompt_tokens = seq_len as u32;
- let mut seqlen_offset = 0;
- let mut generate: Vec = Vec::new();
+ let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
+ let input_features = input_features.to_dtype(self.dtype)?;
+ let audio_token_lengths = Tensor::new(audio_token_lengths, &self.device)?;
let sample_len = mes.max_tokens.unwrap_or(1024);
- for _ in 0..sample_len {
- let logits = self.glm_asr_nano.forward(
- input_features.as_ref(),
- audio_token_lengths.as_ref(),
- &input_ids,
- seqlen_offset,
- )?;
- let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
- let next_token = logit_processor.sample(&logits)?;
- generate.push(next_token);
- 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;
- }
- 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)
+ let mut ctx = GenerationContext::new(
+ mes.temperature,
+ mes.top_p,
+ mes.top_k,
+ seed,
+ input_ids.dim(1)?,
+ sample_len,
+ self.device.clone(),
+ );
+
+ let data_vec = vec![input_features.into(), audio_token_lengths.into()];
+ let data = MultiModalData::new(data_vec);
+ generate_generic(
+ &mut self.glm_asr_nano,
+ &self.tokenizer,
+ input_ids,
+ data,
+ &mut ctx,
+ &self.model_name,
+ )
}
fn generate_stream(
@@ -126,58 +107,29 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'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 render_text = self.chat_template.apply_chat_template(&mes)?;
let (input_features, audio_token_lengths, replace_text) =
self.processor.process_info(&mes, &render_text)?;
let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
-
- let mut seq_len = input_ids.dim(1)?;
- let mut seqlen_offset = 0;
+ let input_features = input_features.to_dtype(self.dtype)?;
+ let audio_token_lengths = Tensor::new(audio_token_lengths, &self.device)?;
let sample_len = mes.max_tokens.unwrap_or(1024);
- let stream = stream! {
- let mut error_tokens = Vec::new();
- let mut input_features = Some(input_features.to_dtype(self.dtype)?);
- let mut audio_token_lengths = Some(audio_token_lengths);
- let mut input_ids = input_ids;
- for _ in 0..sample_len {
- let logits =
- self.glm_asr_nano
- .forward(input_features.as_ref(), audio_token_lengths.as_ref(), &input_ids, 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)?;
- 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();
- };
+ let data_vec = vec![input_features.into(), audio_token_lengths.into()];
+ let data = MultiModalData::new(data_vec);
+ let stream = generate_stream_generic(
+ &mut self.glm_asr_nano,
+ &self.tokenizer,
+ input_ids,
+ data,
+ mes.temperature,
+ mes.top_p,
+ None,
+ seed,
+ sample_len,
+ false,
+ &self.device,
+ &self.model_name,
+ )?;
Ok(Box::new(Box::pin(stream)))
}
}
diff --git a/src/models/glm_asr_nano/model.rs b/src/models/glm_asr_nano/model.rs
index 67d7285..e6d0a68 100644
--- a/src/models/glm_asr_nano/model.rs
+++ b/src/models/glm_asr_nano/model.rs
@@ -4,8 +4,11 @@ use candle_nn::{Conv1d, LayerNorm, Linear, Module, VarBuilder, linear, linear_no
use crate::{
models::{
- common::modules::{
- LlamaForCausalLM, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm,
+ common::{
+ InferenceModel,
+ modules::{
+ LlamaForCausalLM, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm,
+ },
},
glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig},
},
@@ -233,10 +236,11 @@ pub struct GlmAsrNanoModel {
audio_tower: GlmAsrEncoder,
multi_modal_projector: TwoLinearMLP,
language_model: LlamaForCausalLM,
+ stop_token_ids: Vec,
}
impl GlmAsrNanoModel {
- pub fn new(vb: VarBuilder, config: GlmAsrNanoConfig) -> Result {
+ pub fn new(vb: VarBuilder, config: GlmAsrNanoConfig, eos_ids: Vec) -> Result {
let audio_tower = GlmAsrEncoder::new(vb.pp("audio_tower"), &config.audio_config)?;
let multi_modal_projector = TwoLinearMLP::new(
vb.pp("multi_modal_projector"),
@@ -273,6 +277,7 @@ impl GlmAsrNanoModel {
audio_tower,
multi_modal_projector,
language_model,
+ stop_token_ids: eos_ids,
})
}
@@ -300,7 +305,7 @@ impl GlmAsrNanoModel {
pub fn forward(
&mut self,
input_features: Option<&Tensor>,
- audio_token_lengths: Option<&Vec>,
+ audio_token_lengths: Option<&Tensor>,
input_ids: &Tensor,
seqlen_offset: usize,
) -> Result {
@@ -308,8 +313,9 @@ impl GlmAsrNanoModel {
if let Some(input_features) = input_features
&& let Some(audio_token_len) = audio_token_lengths
{
+ let audio_token_len = audio_token_len.to_vec1::()?;
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)?;
}
let logits = self.language_model.forward(&inputs_embeds, seqlen_offset)?;
@@ -319,3 +325,38 @@ impl GlmAsrNanoModel {
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 {
+ 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 {
+ self.forward(None, None, input_ids, seqlen_offset)
+ }
+
+ fn clear_cache(&mut self) {
+ self.clear_kv_cache();
+ }
+
+ fn stop_token_ids(&self) -> Vec {
+ self.stop_token_ids.clone()
+ }
+}
diff --git a/src/models/glm_asr_nano/processor.rs b/src/models/glm_asr_nano/processor.rs
index 7b3f818..184e198 100644
--- a/src/models/glm_asr_nano/processor.rs
+++ b/src/models/glm_asr_nano/processor.rs
@@ -206,6 +206,9 @@ impl GlmAsrNanoProcessor {
render_text: &str,
) -> Result<(Tensor, Vec, String)> {
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) =
self.process_audio(audio_tensors)?;
let audio_lengths = input_features_mask.sum(D::Minus1)?;
diff --git a/src/models/glm_ocr/generate.rs b/src/models/glm_ocr/generate.rs
index caad0e2..e3b3c94 100644
--- a/src/models/glm_ocr/generate.rs
+++ b/src/models/glm_ocr/generate.rs
@@ -1,12 +1,14 @@
//! GLM-OCR Inference and Generation
use crate::{
- models::common::generate::get_logit_processor,
+ models::common::{
+ MultiModalData,
+ generate::{GenerationContext, generate_generic, generate_stream_generic},
+ },
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
};
use anyhow::{Result, anyhow};
-use candle_core::{DType, Device, IndexOp, Tensor};
+use candle_core::{DType, Device};
use candle_nn::VarBuilder;
-use rocket::async_stream::stream;
use rocket::futures::Stream;
use crate::{
@@ -21,8 +23,7 @@ use crate::{
},
tokenizer::TokenizerModel,
utils::{
- build_completion_chunk_response, build_completion_response, extract_user_text,
- find_type_files, get_device, get_dtype, img_utils::extract_image_url,
+ extract_user_text, find_type_files, get_device, get_dtype, img_utils::extract_image_url,
},
};
@@ -32,7 +33,6 @@ pub struct GlmOcrGenerateModel {
processor: GlmOcrProcessor,
model: GlmOcrModel,
device: Device,
- eos_token_ids: Vec,
model_name: String,
image_token_id: u32,
image_start_token_id: u32,
@@ -54,10 +54,10 @@ impl GlmOcrGenerateModel {
let processor = GlmOcrProcessor::new(path, &device, dtype)?;
let model_list = find_type_files(path, "safetensors")?;
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: GlmOcrGenerationConfig =
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)
.file_name()
.and_then(|s| s.to_str())
@@ -69,7 +69,6 @@ impl GlmOcrGenerateModel {
processor,
model,
device,
- eos_token_ids: generation_config.eos_token_id.clone(),
model_name,
image_token_id: cfg.image_token_id,
image_start_token_id: cfg.image_start_token_id,
@@ -84,8 +83,6 @@ impl GlmOcrGenerateModel {
impl GenerateModel for GlmOcrGenerateModel {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result {
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
let image_urls = extract_image_url(&mes);
let image_path = image_urls
@@ -110,55 +107,32 @@ impl GenerateModel for GlmOcrGenerateModel {
self.spatial_merge_size,
)?;
- let mut input_ids = processed.input_ids;
- let pixel_values = Some(processed.pixel_values);
- let image_grid_thw = Some(processed.grid_thw);
- let image_mask = Some(processed.image_mask);
- let mut seqlen_offset = 0;
- let mut seq_len = input_ids.dim(1)?;
- let prompt_tokens = seq_len as u32;
- let mut generate = Vec::new();
- let sample_len = mes.max_tokens.unwrap_or(512);
+ let input_ids = processed.input_ids;
+ let sample_len = mes.max_tokens.unwrap_or(1024);
+ let mut ctx = GenerationContext::new(
+ mes.temperature,
+ mes.top_p,
+ mes.top_k,
+ seed,
+ input_ids.dim(1)?,
+ sample_len,
+ self.device.clone(),
+ );
- for _ in 0..sample_len {
- let is_first_pass = seqlen_offset == 0;
- let logits = self.model.forward(
- &input_ids,
- if is_first_pass {
- pixel_values.as_ref()
- } else {
- None
- },
- if is_first_pass {
- image_grid_thw.as_ref()
- } else {
- None
- },
- 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)
+ let data_vec = vec![
+ processed.pixel_values.into(),
+ processed.grid_thw.into(),
+ processed.image_mask.into(),
+ ];
+ let data = MultiModalData::new(data_vec);
+ generate_generic(
+ &mut self.model,
+ &self.tokenizer,
+ input_ids,
+ data,
+ &mut ctx,
+ &self.model_name,
+ )
}
fn generate_stream(
@@ -173,8 +147,6 @@ impl GenerateModel for GlmOcrGenerateModel {
>,
> {
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
let image_urls = extract_image_url(&mes);
let image_path = image_urls
@@ -199,63 +171,28 @@ impl GenerateModel for GlmOcrGenerateModel {
self.spatial_merge_size,
)?;
- let mut input_ids = processed.input_ids;
- let pixel_values = Some(processed.pixel_values);
- let image_grid_thw = Some(processed.grid_thw);
- let image_mask = Some(processed.image_mask);
- let mut seqlen_offset = 0;
- let mut seq_len = input_ids.dim(1)?;
- let sample_len = mes.max_tokens.unwrap_or(512);
-
- let stream = stream! {
- let mut generated: Vec = Vec::new();
- let mut error_tokens = Vec::new();
- for _ in 0..sample_len {
- let is_first_pass = seqlen_offset == 0;
- let logits = self.model.forward(
- &input_ids,
- if is_first_pass { pixel_values.as_ref() } else { None },
- if is_first_pass { image_grid_thw.as_ref() } else { None },
- if is_first_pass { image_mask.as_ref() } else { None },
- seqlen_offset,
- ).map_err(|e| anyhow!(format!("forward error: {e}")))?;
- 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}")))?;
-
- 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();
- };
-
+ let input_ids = processed.input_ids;
+ let sample_len = mes.max_tokens.unwrap_or(1024);
+ let data_vec = vec![
+ processed.pixel_values.into(),
+ processed.grid_thw.into(),
+ processed.image_mask.into(),
+ ];
+ let data = MultiModalData::new(data_vec);
+ let stream = generate_stream_generic(
+ &mut self.model,
+ &self.tokenizer,
+ input_ids,
+ data,
+ mes.temperature,
+ mes.top_p,
+ None,
+ seed,
+ sample_len,
+ false,
+ &self.device,
+ &self.model_name,
+ )?;
Ok(Box::new(Box::pin(stream)))
}
}
diff --git a/src/models/glm_ocr/model.rs b/src/models/glm_ocr/model.rs
index 912084f..a650461 100644
--- a/src/models/glm_ocr/model.rs
+++ b/src/models/glm_ocr/model.rs
@@ -9,7 +9,7 @@ use candle_nn::{
use crate::{
models::{
- common::modules::GateUpDownMLP,
+ common::{InferenceModel, modules::GateUpDownMLP},
glm_ocr::config::{GlmOcrConfig, GlmOcrTextConfig, GlmOcrVisionConfig},
},
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)?;
- 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)
}
@@ -1271,10 +1272,11 @@ impl GlmOcrTextModel {
pub struct GlmOcrModel {
vision_encoder: GlmOcrVisionModel,
language_model: GlmOcrTextModel,
+ stop_token_ids: Vec,
}
impl GlmOcrModel {
- pub fn new(vb: VarBuilder, config: GlmOcrConfig) -> Result {
+ pub fn new(vb: VarBuilder, config: GlmOcrConfig, eos_ids: Vec) -> Result {
let vision_encoder =
GlmOcrVisionModel::new(vb.pp("model").pp("visual"), &config.vision_config)?;
let language_model = GlmOcrTextModel::new(
@@ -1286,6 +1288,7 @@ impl GlmOcrModel {
Ok(Self {
vision_encoder,
language_model,
+ stop_token_ids: eos_ids,
})
}
@@ -1331,3 +1334,40 @@ impl GlmOcrModel {
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 {
+ 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 {
+ self.forward(input_ids, None, None, None, seqlen_offset)
+ }
+
+ fn clear_cache(&mut self) {
+ self.clear_kv_cache();
+ }
+
+ fn stop_token_ids(&self) -> Vec {
+ self.stop_token_ids.clone()
+ }
+}
diff --git a/src/models/hunyuan_ocr/config.rs b/src/models/hunyuan_ocr/config.rs
index 7cc172c..46bcbb9 100644
--- a/src/models/hunyuan_ocr/config.rs
+++ b/src/models/hunyuan_ocr/config.rs
@@ -85,7 +85,7 @@ pub struct HunyuanOCRGenerationConfig {
pub bos_token_id: usize,
pub pad_token_id: usize,
pub do_sample: bool,
- pub eos_token_id: Vec,
+ pub eos_token_id: Vec,
pub top_p: f32,
pub top_k: usize,
pub temperature: f32,
diff --git a/src/models/hunyuan_ocr/generate.rs b/src/models/hunyuan_ocr/generate.rs
index 1a993d1..8b0e212 100644
--- a/src/models/hunyuan_ocr/generate.rs
+++ b/src/models/hunyuan_ocr/generate.rs
@@ -1,11 +1,13 @@
use crate::{
- models::common::generate::get_logit_processor,
+ models::common::{
+ MultiModalData,
+ generate::{GenerationContext, generate_generic, generate_stream_generic},
+ },
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
};
-use anyhow::{Result, anyhow};
-use candle_core::{DType, Device, Tensor};
+use anyhow::Result;
+use candle_core::{DType, Device};
use candle_nn::VarBuilder;
-use rocket::async_stream::stream;
use rocket::futures::Stream;
use crate::{
@@ -19,10 +21,7 @@ use crate::{
},
},
tokenizer::TokenizerModel,
- utils::{
- build_completion_chunk_response, build_completion_response, find_type_files, get_device,
- get_dtype,
- },
+ utils::{find_type_files, get_device, get_dtype},
};
pub struct HunyuanOCRGenerateModel<'a> {
@@ -31,8 +30,6 @@ pub struct HunyuanOCRGenerateModel<'a> {
pre_processor: HunyuanVLProcessor,
hunyuan_vl: HunyuanVLModel,
device: Device,
- eos_token_id1: u32,
- eos_token_id2: u32,
generation_config: HunyuanOCRGenerationConfig,
model_name: String,
}
@@ -49,10 +46,12 @@ impl<'a> HunyuanOCRGenerateModel<'a> {
let pre_processor = HunyuanVLProcessor::new(path, &device, dtype)?;
let model_list = find_type_files(path, "safetensors")?;
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: HunyuanOCRGenerationConfig =
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)
.file_name()
.and_then(|s| s.to_str())
@@ -64,8 +63,6 @@ impl<'a> HunyuanOCRGenerateModel<'a> {
pre_processor,
hunyuan_vl,
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,
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_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 data = self
.pre_processor
.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 = Vec::new();
let sample_len = mes.max_tokens.unwrap_or(1024);
- for _ in 0..sample_len {
- let logits = self.hunyuan_vl.forward(
- &input_ids,
- pixel_values.as_ref(),
- image_grid_thw.as_ref(),
- image_mask,
- position_ids,
- seqlen_offset,
- )?;
- let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
- let next_token = logit_processor.sample(&logits)?;
- generate.push(next_token);
- 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;
- }
- let num_token = generate.len() as u32;
- let res = self.tokenizer.token_decode(generate)?;
- self.hunyuan_vl.clear_kv_cache();
- let response =
- build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
- Ok(response)
+ let input_ids = data.input_ids;
+ let mut ctx = GenerationContext::new(
+ temperature.into(),
+ top_p.into(),
+ top_k.into(),
+ seed,
+ input_ids.dim(1)?,
+ sample_len,
+ self.device.clone(),
+ );
+
+ let data_vec = vec![
+ data.pixel_values,
+ data.image_grid_thw,
+ data.image_mask.into(),
+ data.position_ids.into(),
+ ];
+ let data = MultiModalData::new(data_vec);
+ generate_generic(
+ &mut self.hunyuan_vl,
+ &self.tokenizer,
+ input_ids,
+ data,
+ &mut ctx,
+ &self.model_name,
+ )
}
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_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 data = self
.pre_processor
.process_info(&mes, &self.tokenizer, &mes_render)?;
- 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 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)?;
- for _ in 0..sample_len {
- let logits = self.hunyuan_vl.forward(
- &input_ids,
- pixel_values.as_ref(),
- image_grid_thw.as_ref(),
- image_mask,
- position_ids,
- 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)?;
- 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();
- };
+ let input_ids = data.input_ids;
+ let data_vec = vec![
+ data.pixel_values,
+ data.image_grid_thw,
+ data.image_mask.into(),
+ data.position_ids.into(),
+ ];
+ let data = MultiModalData::new(data_vec);
+ let stream = generate_stream_generic(
+ &mut self.hunyuan_vl,
+ &self.tokenizer,
+ input_ids,
+ data,
+ temperature.into(),
+ top_p.into(),
+ top_k.into(),
+ seed,
+ sample_len,
+ false,
+ &self.device,
+ &self.model_name,
+ )?;
Ok(Box::new(Box::pin(stream)))
}
}
diff --git a/src/models/hunyuan_ocr/model.rs b/src/models/hunyuan_ocr/model.rs
index e98836a..2585618 100644
--- a/src/models/hunyuan_ocr/model.rs
+++ b/src/models/hunyuan_ocr/model.rs
@@ -7,14 +7,19 @@ use candle_nn::{
use crate::{
models::{
- common::modules::{
- GateUpDownMLP, NaiveAttnTwoLinearMLPBlock, eager_attention_forward, get_conv2d,
+ common::{
+ InferenceModel,
+ modules::{
+ GateUpDownMLP, NaiveAttnTwoLinearMLPBlock, eager_attention_forward, get_conv2d,
+ },
},
hunyuan_ocr::config::{HunYuanVLConfig, HunYuanVLVisionConfig},
},
position_embed::rope::{RoPE, apply_rotary_pos_emb, get_xd_cos_sin},
- utils::interpolate::interpolate_bilinear,
- utils::tensor_utils::{masked_scatter_dim0, prepare_causal_attention_mask, split_tensor},
+ utils::{
+ interpolate::interpolate_bilinear,
+ tensor_utils::{masked_scatter_dim0, prepare_causal_attention_mask, split_tensor},
+ },
};
pub struct HunYuanVisionPatchEmbed {
@@ -538,10 +543,11 @@ pub struct HunyuanVLModel {
vit: HunYuanVisionTransformer,
model: HunYuanVLTextModel,
lm_head: Linear,
+ stop_token_ids: Vec,
}
impl HunyuanVLModel {
- pub fn new(vb: VarBuilder, config: HunYuanVLConfig) -> Result {
+ pub fn new(vb: VarBuilder, config: HunYuanVLConfig, eos_ids: Vec) -> Result {
let vit = HunYuanVisionTransformer::new(vb.pp("vit"), &config.vision_config)?;
let model = HunYuanVLTextModel::new(vb.pp("model"), &config)?;
let lm_head = Linear::new(model.embed_tokens.embeddings().clone(), None);
@@ -550,6 +556,7 @@ impl HunyuanVLModel {
vit,
model,
lm_head,
+ stop_token_ids: eos_ids,
})
}
pub fn forward(
@@ -582,3 +589,42 @@ impl HunyuanVLModel {
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 {
+ 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 {
+ 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 {
+ self.stop_token_ids.clone()
+ }
+}
diff --git a/src/models/lfm2/generate.rs b/src/models/lfm2/generate.rs
index 202fd7b..d81d794 100644
--- a/src/models/lfm2/generate.rs
+++ b/src/models/lfm2/generate.rs
@@ -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::utils::build_completion_chunk_response;
use crate::{
chat_template::ChatTemplate,
models::{
@@ -11,19 +13,17 @@ use crate::{
},
},
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 candle_core::{DType, Device, Tensor};
+use candle_core::{DType, Device};
use candle_nn::VarBuilder;
-use rocket::async_stream::stream;
pub struct Lfm2GenerateModel<'a> {
chat_template: ChatTemplate<'a>,
tokenizer: TokenizerModel,
device: Device,
model: Lfm2Model,
- eos_token_id: u32,
model_name: String,
}
impl<'a> Lfm2GenerateModel<'a> {
@@ -45,8 +45,8 @@ impl<'a> Lfm2GenerateModel<'a> {
};
let dtype = get_dtype(dtype, &cfg_dtype);
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_path, dtype, &device)? };
- let model = Lfm2Model::new(vb, &cfg)?;
- let eos_token_id = gen_cfg.eos_token_id;
+ let eos_ids = vec![gen_cfg.eos_token_id];
+ let model = Lfm2Model::new(vb, &cfg, eos_ids)?;
let model_name = std::path::Path::new(path)
.file_name()
.and_then(|s| s.to_str())
@@ -57,7 +57,6 @@ impl<'a> Lfm2GenerateModel<'a> {
tokenizer,
device,
model,
- eos_token_id,
model_name,
})
}
@@ -66,40 +65,28 @@ impl<'a> Lfm2GenerateModel<'a> {
impl<'a> GenerateModel for Lfm2GenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result {
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.top_p,
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 prompt_tokens = seq_len as u32;
- let mut seqlen_offset = 0;
- let mut generate = vec![];
- let sample_len = mes.max_tokens.unwrap_or(1024);
- 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)?;
- 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,
+
+ let data = MultiModalData::new(vec![]);
+ generate_generic(
+ &mut self.model,
+ &self.tokenizer,
+ input_ids,
+ data,
+ &mut ctx,
&self.model_name,
- Some(completion_tokens),
- Some(prompt_tokens),
- );
- Ok(mes)
+ )
}
fn generate_stream(
@@ -115,50 +102,24 @@ impl<'a> GenerateModel for Lfm2GenerateModel<'a> {
>,
> {
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.top_p,
None,
- mes.seed.unwrap_or(34562) as u64,
- );
- 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(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();
- };
+ seed,
+ sample_len,
+ false,
+ &self.device,
+ &self.model_name,
+ )?;
Ok(Box::new(Box::pin(stream)))
}
}
diff --git a/src/models/lfm2/model.rs b/src/models/lfm2/model.rs
index 348b957..e7ad4d3 100644
--- a/src/models/lfm2/model.rs
+++ b/src/models/lfm2/model.rs
@@ -1,6 +1,9 @@
use crate::{
models::{
- common::modules::{GateUpDownMLP, QKNormAttention, conv1d_depthwise, get_conv1d},
+ common::{
+ InferenceModel,
+ modules::{GateUpDownMLP, QKNormAttention, conv1d_depthwise, get_conv1d},
+ },
lfm2::config::Lfm2Config,
},
position_embed::rope::RoPE,
@@ -279,10 +282,11 @@ impl Lfm2Decoder {
pub struct Lfm2Model {
model: Lfm2Decoder,
lm_head: Linear,
+ stop_token_ids: Vec,
}
impl Lfm2Model {
- pub fn new(vb: VarBuilder, config: &Lfm2Config) -> Result {
+ pub fn new(vb: VarBuilder, config: &Lfm2Config, eos_ids: Vec) -> Result {
let model = Lfm2Decoder::new(vb.pp("model"), config)?;
let lm_head = if let Some(flag) = config.tie_embedding
&& flag
@@ -300,7 +304,11 @@ impl Lfm2Model {
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 {
@@ -315,3 +323,17 @@ impl Lfm2Model {
self.model.clear_cache();
}
}
+
+impl InferenceModel for Lfm2Model {
+ fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result {
+ self.forward(input_ids, seqlen_offset)
+ }
+
+ fn clear_cache(&mut self) {
+ self.clear_cache();
+ }
+
+ fn stop_token_ids(&self) -> Vec {
+ self.stop_token_ids.clone()
+ }
+}
diff --git a/src/models/lfm2vl/generate.rs b/src/models/lfm2vl/generate.rs
index db25985..e78c6f3 100644
--- a/src/models/lfm2vl/generate.rs
+++ b/src/models/lfm2vl/generate.rs
@@ -1,9 +1,12 @@
use crate::{
- models::common::generate::get_logit_processor,
+ models::common::{
+ MultiModalData,
+ generate::{GenerationContext, generate_generic, generate_stream_generic},
+ },
params::chat::{ChatCompletionParameters, ChatCompletionResponse},
};
use anyhow::Result;
-use candle_core::{DType, Device, Tensor};
+use candle_core::{DType, Device};
use candle_nn::VarBuilder;
use crate::{
@@ -14,12 +17,8 @@ use crate::{
lfm2vl::{config::Lfm2VLConfig, model::Lfm2VLModel, processor::Lfm2VLProcessor},
},
tokenizer::TokenizerModel,
- utils::{
- build_completion_chunk_response, build_completion_response, find_type_files, get_device,
- get_dtype,
- },
+ utils::{find_type_files, get_device, get_dtype},
};
-use rocket::async_stream::stream;
pub struct Lfm2VLGenerateModel<'a> {
chat_template: ChatTemplate<'a>,
@@ -27,7 +26,6 @@ pub struct Lfm2VLGenerateModel<'a> {
device: Device,
model: Lfm2VLModel,
processor: Lfm2VLProcessor,
- eos_token_id: u32,
model_name: String,
}
impl<'a> Lfm2VLGenerateModel<'a> {
@@ -43,9 +41,9 @@ impl<'a> Lfm2VLGenerateModel<'a> {
let model_path = find_type_files(path, "safetensors")?;
let dtype = get_dtype(dtype, &cfg.dtype);
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 eos_token_id = gen_cfg.eos_token_id;
let model_name = std::path::Path::new(path)
.file_name()
.and_then(|s| s.to_str())
@@ -57,7 +55,6 @@ impl<'a> Lfm2VLGenerateModel<'a> {
device,
model,
processor,
- eos_token_id,
model_name,
})
}
@@ -66,54 +63,35 @@ impl<'a> Lfm2VLGenerateModel<'a> {
impl<'a> GenerateModel for Lfm2VLGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result {
let mes_render = self.chat_template.apply_chat_template(&mes)?;
- 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.top_p,
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 mut input_ids = self.tokenizer.text_encode(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 generate = vec![];
- 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);
- 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)?;
- 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,
+
+ let data_vec = vec![
+ pixel_values.into(),
+ pixel_attention_mask.into(),
+ spatial_shapes.into(),
+ ];
+ let data = MultiModalData::new(data_vec);
+ generate_generic(
+ &mut self.model,
+ &self.tokenizer,
+ input_ids,
+ data,
+ &mut ctx,
&self.model_name,
- Some(completion_tokens),
- Some(prompt_tokens),
- );
- Ok(mes)
+ )
}
fn generate_stream(
@@ -129,67 +107,31 @@ impl<'a> GenerateModel for Lfm2VLGenerateModel<'a> {
>,
> {
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.top_p,
None,
- mes.seed.unwrap_or(34562) as u64,
- );
- let (pixel_values, pixel_attention_mask, spatial_shapes, text) =
- self.processor.process_info(&mes, &mes_render)?;
- let mut input_ids = self.tokenizer.text_encode(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 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();
- };
+ seed,
+ sample_len,
+ false,
+ &self.device,
+ &self.model_name,
+ )?;
Ok(Box::new(Box::pin(stream)))
}
}
diff --git a/src/models/lfm2vl/model.rs b/src/models/lfm2vl/model.rs
index 7fcf394..3365a58 100644
--- a/src/models/lfm2vl/model.rs
+++ b/src/models/lfm2vl/model.rs
@@ -1,6 +1,9 @@
use crate::{
models::{
- common::modules::{NaiveAttnTwoLinearMLPBlock, get_layer_norm},
+ common::{
+ InferenceModel,
+ modules::{NaiveAttnTwoLinearMLPBlock, get_layer_norm},
+ },
lfm2::model::Lfm2Decoder,
lfm2vl::config::{Lfm2VLConfig, Lfm2VLVisionConfig},
},
@@ -15,11 +18,7 @@ use candle_nn::{Activation, LayerNorm, Linear, Module, VarBuilder, embedding, li
use num::integer::Roots;
pub struct Siglip2VisionEmbeddings {
- // embed_dim: usize,
- // patch_size: usize,
patch_embedding: Linear,
- // position_embedding_size: usize,
- // position_embedding: Embedding,
postitional_embeddings: Tensor,
}
@@ -44,11 +43,7 @@ impl Siglip2VisionEmbeddings {
.permute((2, 0, 1))?
.unsqueeze(0)?;
Ok(Self {
- // embed_dim,
- // patch_size,
patch_embedding,
- // position_embedding_size,
- // position_embedding,
postitional_embeddings,
})
}
@@ -254,10 +249,11 @@ pub struct Lfm2VLModel {
language_model: Lfm2Decoder,
lm_head: Linear,
img_id: u32,
+ stop_token_ids: Vec,
}
impl Lfm2VLModel {
- pub fn new(vb: VarBuilder, cfg: &Lfm2VLConfig) -> Result {
+ pub fn new(vb: VarBuilder, cfg: &Lfm2VLConfig, eos_ids: Vec) -> Result {
let vb = vb.pp("model");
let vision_tower = Siglip2VisionModel::new(vb.pp("vision_tower"), &cfg.vision_config)?;
let multi_modal_projector =
@@ -270,6 +266,7 @@ impl Lfm2VLModel {
language_model,
lm_head,
img_id: cfg.image_token_id,
+ stop_token_ids: eos_ids,
})
}
@@ -321,3 +318,40 @@ impl Lfm2VLModel {
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 {
+ 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 {
+ self.forward(input_ids, None, None, None, seqlen_offset)
+ }
+
+ fn clear_cache(&mut self) {
+ self.clear_cache();
+ }
+
+ fn stop_token_ids(&self) -> Vec {
+ self.stop_token_ids.clone()
+ }
+}
diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs
index 58d96a7..d586c34 100644
--- a/src/models/minicpm4/generate.rs
+++ b/src/models/minicpm4/generate.rs
@@ -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::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
-use anyhow::{Result, anyhow};
-use candle_core::{DType, Device, Tensor};
+use anyhow::Result;
+use candle_core::{DType, Device};
use candle_nn::VarBuilder;
-use rocket::async_stream::stream;
use rocket::futures::Stream;
use crate::models::minicpm4::config::MiniCPM4Config;
use crate::models::minicpm4::model::MiniCPMModel;
// use crate::models::GenerateStream;
-use crate::utils::{
- build_completion_chunk_response, build_completion_response, find_type_files, get_device,
- get_dtype,
-};
+use crate::utils::{find_type_files, get_device, get_dtype};
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
pub struct MiniCPMGenerateModel<'a> {
@@ -22,8 +21,6 @@ pub struct MiniCPMGenerateModel<'a> {
tokenizer: TokenizerModel,
minicpm: MiniCPMModel,
device: Device,
- endoftext_id: u32,
- im_end_id: u32,
model_name: String,
}
@@ -40,7 +37,8 @@ impl<'a> MiniCPMGenerateModel<'a> {
let im_end_id = cfg.eos_token_id[1];
let model_list = find_type_files(path, "safetensors")?;
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)
.file_name()
.and_then(|s| s.to_str())
@@ -51,8 +49,6 @@ impl<'a> MiniCPMGenerateModel<'a> {
tokenizer,
minicpm,
device: device.clone(),
- endoftext_id,
- im_end_id,
model_name,
})
}
@@ -60,33 +56,29 @@ impl<'a> MiniCPMGenerateModel<'a> {
impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result {
- 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 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 input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
+ let seed = mes.seed.unwrap_or(34562) as u64;
let sample_len = mes.max_tokens.unwrap_or(2048);
- for _ in 0..sample_len {
- let logits = self.minicpm.forward_with_cache(&input_ids, seqlen_offset)?;
- let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
- let next_token = logit_processor.sample(&logits)?;
- generate.push(next_token);
- 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)?;
- }
- let num_token = generate.len() as u32;
- let res = self.tokenizer.token_decode(generate)?;
- self.minicpm.clear_kv_cache();
- let response =
- build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
- Ok(response)
+ let mut ctx = GenerationContext::new(
+ mes.temperature,
+ mes.top_p,
+ None,
+ seed,
+ input_ids.dim(1)?,
+ sample_len,
+ self.device.clone(),
+ );
+
+ let data = MultiModalData::new(vec![]);
+ generate_generic(
+ &mut self.minicpm,
+ &self.tokenizer,
+ input_ids,
+ data,
+ &mut ctx,
+ &self.model_name,
+ )
}
fn generate_stream(
&mut self,
@@ -100,50 +92,24 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'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 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 input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
+ let data = MultiModalData::new(vec![]);
let sample_len = mes.max_tokens.unwrap_or(512);
- let stream = stream! {
- let mut error_tokens = Vec::new();
- for _ in 0..sample_len {
- let logits = self.minicpm.forward_with_cache(
- &input_ids,
- 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)?;
- 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();
- };
+ let stream = generate_stream_generic(
+ &mut self.minicpm,
+ &self.tokenizer,
+ input_ids,
+ data,
+ mes.temperature,
+ mes.top_p,
+ None,
+ seed,
+ sample_len,
+ false,
+ &self.device,
+ &self.model_name,
+ )?;
Ok(Box::new(Box::pin(stream)))
}
}
diff --git a/src/models/minicpm4/model.rs b/src/models/minicpm4/model.rs
index 2c43686..a536ada 100644
--- a/src/models/minicpm4/model.rs
+++ b/src/models/minicpm4/model.rs
@@ -4,7 +4,10 @@ use candle_nn::{Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, rms_n
use crate::{
models::{
- common::modules::{GateUpDownMLP, NaiveAttention},
+ common::{
+ InferenceModel,
+ modules::{GateUpDownMLP, NaiveAttention},
+ },
minicpm4::config::MiniCPM4Config,
},
position_embed::rope::compute_default_rope_parameters,
@@ -207,10 +210,11 @@ pub struct MiniCPMModel {
norm: RmsNorm,
rope_emb: MiniCPMLongRoPE,
lm_head: Linear,
+ stop_token_ids: Vec,
}
impl MiniCPMModel {
- pub fn new(vb: VarBuilder, cfg: MiniCPM4Config) -> Result {
+ pub fn new(vb: VarBuilder, cfg: MiniCPM4Config, eos_ids: Vec) -> Result {
let vb = vb.pp("model");
let embed_tokens = embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("embed_tokens"))?;
let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
@@ -229,10 +233,11 @@ impl MiniCPMModel {
norm,
rope_emb,
lm_head,
+ stop_token_ids: eos_ids,
})
}
- pub fn forward(&mut self, input_ids: &Tensor, position_id: usize) -> Result {
+ pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result {
let (bs, seq_len) = input_ids.dims2()?;
let input_embeds = self
.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;
for decode_layer in &self.layers {
hidden_states =
@@ -267,7 +272,11 @@ impl MiniCPMModel {
Ok(logits)
}
- pub fn forward_with_cache(&mut self, input_ids: &Tensor, position_id: usize) -> Result {
+ pub fn forward_with_cache(
+ &mut self,
+ input_ids: &Tensor,
+ seqlen_offset: usize,
+ ) -> Result {
let (bs, seq_len) = input_ids.dims2()?;
let input_embeds = self
.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;
for decode_layer in &mut self.layers {
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 {
+ self.forward_with_cache(input_ids, seqlen_offset)
+ }
+
+ fn clear_cache(&mut self) {
+ self.clear_kv_cache();
+ }
+
+ fn stop_token_ids(&self) -> Vec {
+ self.stop_token_ids.clone()
+ }
+}
diff --git a/src/models/paddleocr_vl/generate.rs b/src/models/paddleocr_vl/generate.rs
index 604ff6e..5a39a03 100644
--- a/src/models/paddleocr_vl/generate.rs
+++ b/src/models/paddleocr_vl/generate.rs
@@ -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::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
-use anyhow::{Result, anyhow};
+use anyhow::Result;
use candle_core::{D, DType, Device, IndexOp, Tensor};
use candle_nn::VarBuilder;
-use rocket::async_stream::stream;
use rocket::futures::Stream;
use crate::models::paddleocr_vl::config::{PaddleOCRVLConfig, PaddleOCRVLPreprocessorConfig};
use crate::models::paddleocr_vl::model::PaddleOCRVLModel;
use crate::models::paddleocr_vl::processor::PaddleOCRVLProcessor;
use crate::utils::tensor_utils::get_equal_mask;
-use crate::utils::{
- build_completion_chunk_response, build_completion_response, find_type_files, get_device,
- get_dtype,
-};
+use crate::utils::{find_type_files, get_device, get_dtype};
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
pub struct PaddleOCRVLGenerateModel<'a> {
@@ -25,7 +24,6 @@ pub struct PaddleOCRVLGenerateModel<'a> {
paddleocr_vl: PaddleOCRVLModel,
cfg: PaddleOCRVLConfig,
device: Device,
- end_token_id: u32,
model_name: String,
}
@@ -42,10 +40,9 @@ impl<'a> PaddleOCRVLGenerateModel<'a> {
let processor_cfg: PaddleOCRVLPreprocessorConfig =
serde_json::from_slice(&std::fs::read(processor_cfg_path)?)?;
let pre_processor = PaddleOCRVLProcessor::new(processor_cfg, device, dtype)?;
- let end_token_id = 2;
let model_list = find_type_files(path, "safetensors")?;
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)
.file_name()
.and_then(|s| s.to_str())
@@ -58,7 +55,6 @@ impl<'a> PaddleOCRVLGenerateModel<'a> {
paddleocr_vl,
cfg,
device: device.clone(),
- end_token_id,
model_name,
})
}
@@ -66,53 +62,43 @@ impl<'a> PaddleOCRVLGenerateModel<'a> {
impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result {
- 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 (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)?;
- let mut 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 input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
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)?
.cumsum(D::Minus1)?
.to_dtype(candle_core::DType::U32)?
.broadcast_sub(&Tensor::new(vec![1_u32], input_ids.device())?)?;
-
- let mut generate = Vec::new();
+ let seed = mes.seed.unwrap_or(34562) as u64;
let sample_len = mes.max_tokens.unwrap_or(1024);
- for _ in 0..sample_len {
- let logits = self.paddleocr_vl.forward(
- &input_ids,
- pixel_values.as_ref(),
- image_grid_thw.as_ref(),
- &image_mask,
- Some(&cache_position),
- seqlen_offset,
- )?;
- let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
- let next_token = logit_processor.sample(&logits)?;
- generate.push(next_token);
- 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;
- }
- let num_token = generate.len() as u32;
- 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)
+ let mut ctx = GenerationContext::new(
+ mes.temperature,
+ mes.top_p,
+ None,
+ seed,
+ input_ids.dim(1)?,
+ sample_len,
+ self.device.clone(),
+ );
+ let data_vec = vec![
+ pixel_values,
+ image_grid_thw,
+ image_mask.into(),
+ cache_position.into(),
+ ];
+ let data = MultiModalData::new(data_vec);
+ generate_generic(
+ &mut self.paddleocr_vl,
+ &self.tokenizer,
+ input_ids,
+ data,
+ &mut ctx,
+ &self.model_name,
+ )
}
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 (replace_text, pixel_values, image_grid_thw) =
self.pre_processor.process_info(&mes, &mes_render)?;
- let mut input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
- let mut seq_len = input_ids.dim(1)?;
- let mut seqlen_offset = 0;
+ let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
let image_mask = get_equal_mask(&input_ids, self.cfg.image_token_id)?;
- let 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)?
.cumsum(D::Minus1)?
.to_dtype(candle_core::DType::U32)?
.broadcast_sub(&Tensor::new(vec![1_u32], input_ids.device())?)?;
let sample_len = mes.max_tokens.unwrap_or(1024);
- let stream = stream! {
- 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,
- image_grid_thw,
- &image_mask,
- 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;
- 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();
- };
+ let data_vec = vec![
+ pixel_values,
+ image_grid_thw,
+ image_mask.into(),
+ cache_position.into(),
+ ];
+ let data = MultiModalData::new(data_vec);
+ let seed = mes.seed.unwrap_or(34562) as u64;
+ let stream = generate_stream_generic(
+ &mut self.paddleocr_vl,
+ &self.tokenizer,
+ input_ids,
+ data,
+ mes.temperature,
+ mes.top_p,
+ None,
+ seed,
+ sample_len,
+ false,
+ &self.device,
+ &self.model_name,
+ )?;
Ok(Box::new(Box::pin(stream)))
}
}
diff --git a/src/models/paddleocr_vl/model.rs b/src/models/paddleocr_vl/model.rs
index 2305b2c..76eefae 100644
--- a/src/models/paddleocr_vl/model.rs
+++ b/src/models/paddleocr_vl/model.rs
@@ -8,18 +8,23 @@ use num::integer::Roots;
use crate::{
models::{
- common::modules::{
- NaiveAttnGateUpDownMLPBlock, NaiveAttnTwoLinearMLPBlock, get_conv2d, get_layer_norm,
+ common::{
+ InferenceModel,
+ modules::{
+ NaiveAttnGateUpDownMLPBlock, NaiveAttnTwoLinearMLPBlock, get_conv2d, get_layer_norm,
+ },
},
paddleocr_vl::config::{
PaddleOCRVLConfig, PaddleOCRVLRopeScalingConfig, PaddleOCRVLVisionConfig,
},
},
position_embed::rope::{Qwen2_5VLTextRotaryEmbedding, Qwen2_5VisionRotaryEmbedding},
- utils::interpolate::interpolate_bilinear,
- utils::tensor_utils::{
- get_vision_next_indices, masked_scatter_dim0, nonzero_index, prepare_causal_attention_mask,
- zero_index,
+ utils::{
+ interpolate::interpolate_bilinear,
+ tensor_utils::{
+ 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,
lm_head: Linear,
rope_deltas: Option,
+ stop_token_ids: Vec,
}
impl PaddleOCRVLModel {
- pub fn new(cfg: PaddleOCRVLConfig, vb: VarBuilder) -> Result {
+ pub fn new(cfg: PaddleOCRVLConfig, vb: VarBuilder, eos_ids: Vec) -> Result {
let mlp_ar = Projector::new(vb.pp("mlp_AR"), &cfg)?;
let visual = SiglipVisionModel::new(vb.pp("visual"), &cfg.vision_config)?;
let model = Ernie4_5Model::new(vb.pp("model"), &cfg)?;
@@ -433,6 +439,7 @@ impl PaddleOCRVLModel {
cfg,
lm_head,
rope_deltas: None,
+ stop_token_ids: eos_ids,
})
}
@@ -662,13 +669,14 @@ impl PaddleOCRVLModel {
input_ids: &Tensor,
pixel_values: Option<&Tensor>,
image_grid_thw: Option<&Tensor>,
- image_mask: &Tensor,
+ image_mask: Option<&Tensor>,
cache_position: Option<&Tensor>,
seqlen_offset: usize,
) -> Result {
let mut inputs_embeds = self.model.embed_tokens.forward(input_ids)?;
if let Some(pixel_values) = pixel_values
&& let Some(image_grid_thw) = image_grid_thw
+ && let Some(image_mask) = image_mask
{
let pixel_values = pixel_values.unsqueeze(0)?;
let mut siglip_position_ids = vec![];
@@ -716,6 +724,15 @@ impl PaddleOCRVLModel {
.broadcast_add(rope_deltas)?
.contiguous()?
.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 {
Tensor::zeros(1, inputs_embeds.dtype(), inputs_embeds.device())?
};
@@ -740,3 +757,42 @@ impl PaddleOCRVLModel {
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 {
+ 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 {
+ 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 {
+ self.stop_token_ids.clone()
+ }
+}
diff --git a/src/models/qwen2_5vl/generate.rs b/src/models/qwen2_5vl/generate.rs
index 94531dd..3203c09 100644
--- a/src/models/qwen2_5vl/generate.rs
+++ b/src/models/qwen2_5vl/generate.rs
@@ -1,7 +1,12 @@
+use std::time::Instant;
+
use crate::models::common::generate::get_logit_processor;
use crate::params::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
+use crate::utils::response_utils::{
+ build_chunk_response_with_usage, build_completion_response_with_time,
+};
use anyhow::{Result, anyhow};
use candle_core::{D, DType, Device, IndexOp, Tensor};
use candle_nn::VarBuilder;
@@ -10,8 +15,7 @@ use rocket::futures::Stream;
use crate::models::qwen2_5vl::config::Qwen2_5VLConfig;
use crate::utils::{
- build_completion_chunk_response, build_completion_response, find_type_files, get_device,
- get_dtype,
+ find_type_files, get_device, get_dtype, response_utils::build_completion_chunk_response,
};
use crate::{
chat_template::ChatTemplate,
@@ -94,7 +98,10 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
let mut generate = Vec::new();
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 {
+ let i_start = Instant::now();
let logits = self.qwen2_5_vl.forward(
&input_ids,
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 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);
if next_token == self.endoftext_id || next_token == self.im_end_id {
break;
@@ -124,8 +137,14 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?;
self.qwen2_5_vl.clear_kv_cache();
- let response =
- build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
+ let response = build_completion_response_with_time(
+ res,
+ &self.model_name,
+ num_token.into(),
+ completion_secs.into(),
+ prompt_tokens.into(),
+ prompt_secs.into(),
+ );
Ok(response)
}
@@ -148,6 +167,10 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
.tokenizer
.text_encode(input.replace_text.clone(), &self.device)?;
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 mask = Tensor::ones_like(&input_ids)?;
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_content = String::new();
for _ in 0..sample_len {
+ let i_start = Instant::now();
let logits = self.qwen2_5_vl.forward(
&input_ids,
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 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();
if !error_tokens.is_empty() {
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 {
+ yield Ok(build_chunk_response_with_usage(&self.model_name, completion_tokens.into(), completion_secs.into(), prompt_tokens.into(), prompt_secs.into()));
break;
}
seqlen_offset += seq_len;
diff --git a/src/models/qwen3/generate.rs b/src/models/qwen3/generate.rs
index ffe7b46..5c65191 100644
--- a/src/models/qwen3/generate.rs
+++ b/src/models/qwen3/generate.rs
@@ -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::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
-use anyhow::{Result, anyhow};
-use candle_core::{DType, Device, Tensor};
+use anyhow::Result;
+use candle_core::{DType, Device};
use candle_nn::VarBuilder;
-use rocket::async_stream::stream;
use rocket::futures::Stream;
use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig};
use crate::models::qwen3::model::Qwen3Model;
-// use crate::models::GenerateStream;
-use crate::utils::{
- build_completion_chunk_response, build_completion_response, find_type_files, get_device,
- get_dtype,
-};
+use crate::utils::{find_type_files, get_device, get_dtype};
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
pub struct Qwen3GenerateModel<'a> {
@@ -22,8 +20,6 @@ pub struct Qwen3GenerateModel<'a> {
tokenizer: TokenizerModel,
qwen3: Qwen3Model,
device: Device,
- eos_token_id1: u32,
- eos_token_id2: u32,
generation_config: Qwen3GenerationConfig,
model_name: String,
}
@@ -39,10 +35,11 @@ impl<'a> Qwen3GenerateModel<'a> {
let dtype = get_dtype(dtype, cfg_dtype);
let model_list = find_type_files(path, "safetensors")?;
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: Qwen3GenerationConfig =
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)
.file_name()
.and_then(|s| s.to_str())
@@ -53,8 +50,6 @@ impl<'a> Qwen3GenerateModel<'a> {
tokenizer,
qwen3,
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,
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_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 enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
- // 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 input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let sample_len = mes.max_tokens.unwrap_or(2048);
- for _ in 0..sample_len {
- let logits = self.qwen3.forward(Some(&input_ids), None, seqlen_offset)?;
- let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
- let next_token = logit_processor.sample(&logits)?;
- generate.push(next_token);
- 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)?;
- }
- let num_token = generate.len() as u32;
- let res = self.tokenizer.token_decode(generate)?;
- self.qwen3.clear_kv_cache();
- let response =
- build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
- Ok(response)
+ let mut ctx = GenerationContext::new(
+ temperature.into(),
+ top_p.into(),
+ top_k.into(),
+ seed,
+ input_ids.dim(1)?,
+ sample_len,
+ self.device.clone(),
+ );
+
+ let data = MultiModalData::new(vec![]);
+ generate_generic(
+ &mut self.qwen3,
+ &self.tokenizer,
+ input_ids,
+ data,
+ &mut ctx,
+ &self.model_name,
+ )
}
fn generate_stream(
&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_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 enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
- // 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 mut seqlen_offset = 0;
+ let in_reasoning = mes_render.ends_with("\n");
+ let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
+ let data = MultiModalData::new(vec![]);
let sample_len = mes.max_tokens.unwrap_or(512);
- let stream = stream! {
- let mut error_tokens = Vec::new();
- for _ in 0..sample_len {
- let logits = self.qwen3.forward(
- Some(&input_ids),
- None,
- 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)?;
- 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();
- };
+ let stream = generate_stream_generic(
+ &mut self.qwen3,
+ &self.tokenizer,
+ input_ids,
+ data,
+ temperature.into(),
+ top_p.into(),
+ top_k.into(),
+ seed,
+ sample_len,
+ in_reasoning,
+ &self.device,
+ &self.model_name,
+ )?;
Ok(Box::new(Box::pin(stream)))
}
}
diff --git a/src/models/qwen3/model.rs b/src/models/qwen3/model.rs
index cb6ee0d..ef82def 100644
--- a/src/models/qwen3/model.rs
+++ b/src/models/qwen3/model.rs
@@ -6,7 +6,10 @@ use candle_nn::{
use crate::{
models::{
- common::modules::{GateUpDownMLP, QKNormAttention},
+ common::{
+ InferenceModel,
+ modules::{GateUpDownMLP, QKNormAttention},
+ },
qwen3::config::Qwen3Config,
},
position_embed::rope::RoPE,
@@ -94,10 +97,11 @@ pub struct Qwen3Model {
norm: RmsNorm,
rotary_emb: RoPE,
lm_head: Linear,
+ stop_token_ids: Vec,
}
impl Qwen3Model {
- pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result {
+ pub fn new(config: &Qwen3Config, vb: VarBuilder, eos_ids: Vec) -> Result {
let vb = vb.pp("model");
let vocab_size = config.vocab_size;
let embed_tokens = embedding(vocab_size, config.hidden_size, vb.pp("embed_tokens"))?;
@@ -121,6 +125,7 @@ impl Qwen3Model {
norm,
rotary_emb,
lm_head,
+ stop_token_ids: eos_ids,
})
}
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 {
+ self.forward(input_ids.into(), None, seqlen_offset)
+ }
+
+ fn clear_cache(&mut self) {
+ self.clear_kv_cache();
+ }
+
+ fn stop_token_ids(&self) -> Vec {
+ self.stop_token_ids.clone()
+ }
+}
diff --git a/src/models/qwen3_5/generate.rs b/src/models/qwen3_5/generate.rs
index 58da587..25dc8d2 100644
--- a/src/models/qwen3_5/generate.rs
+++ b/src/models/qwen3_5/generate.rs
@@ -1,11 +1,13 @@
use crate::{
- models::common::generate::get_logit_processor,
+ models::common::{
+ MultiModalData,
+ generate::{GenerationContext, generate_generic, generate_stream_generic},
+ },
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
};
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 rocket::async_stream::stream;
use rocket::futures::Stream;
use crate::{
@@ -17,10 +19,7 @@ use crate::{
qwen3vl::processor::Qwen3VLProcessor,
},
tokenizer::TokenizerModel,
- utils::{
- build_completion_chunk_response, build_completion_response, find_type_files, get_device,
- get_dtype,
- },
+ utils::{find_type_files, get_device, get_dtype},
};
pub struct Qwen3_5GenerateModel<'a> {
@@ -29,10 +28,9 @@ pub struct Qwen3_5GenerateModel<'a> {
pre_processor: Option,
qwen3_5: Qwen3_5Model,
device: Device,
- eos_token_id: u32,
model_name: String,
- repeat_penalty: f32,
- repeat_last_n: usize,
+ // repeat_penalty: f32, // TODO
+ // repeat_last_n: usize,
}
impl<'a> Qwen3_5GenerateModel<'a> {
@@ -51,8 +49,8 @@ impl<'a> Qwen3_5GenerateModel<'a> {
let pre_processor = Qwen3VLProcessor::new(path, &device, dtype)?;
let model_list = find_type_files(path, "safetensors")?;
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
- let eos_token_id = cfg.text_config.eos_token_id;
- let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg)?;
+ let eos_ids = vec![cfg.text_config.eos_token_id];
+ let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg, eos_ids)?;
Ok(Self {
chat_template,
@@ -60,10 +58,9 @@ impl<'a> Qwen3_5GenerateModel<'a> {
pre_processor: Some(pre_processor),
qwen3_5,
device,
- eos_token_id,
model_name: model_name.to_string(),
- repeat_penalty: 1.01,
- repeat_last_n: 64,
+ // repeat_penalty: 1.01,
+ // repeat_last_n: 64,
})
}
@@ -105,7 +102,9 @@ impl<'a> Qwen3_5GenerateModel<'a> {
let eos_token_id = model_gguf
.get_matedata("tokenizer.ggml.eos_token_id")?
.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)
.file_stem() // 获取文件名主干(不含扩展名)
.and_then(|s| s.to_str())
@@ -116,11 +115,9 @@ impl<'a> Qwen3_5GenerateModel<'a> {
pre_processor,
qwen3_5,
device,
- // eos_token_id: 248044,
- eos_token_id,
model_name: stem.to_string(),
- repeat_penalty: 1.1,
- repeat_last_n: 64,
+ // repeat_penalty: 1.1,
+ // 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 temperature = mes.temperature.unwrap_or(0.4);
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_text, pixel_values, image_grid_thw, pixel_values_video, video_grid_thw) =
if let Some(processor) = &self.pre_processor {
@@ -147,57 +141,32 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
} else {
(mes_render, None, None, None, None)
};
- let mut 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 input_ids = self.tokenizer.text_encode(mes_text, &self.device)?;
let sample_len = mes.max_tokens.unwrap_or(1024);
- for _ in 0..sample_len {
- let logits = self.qwen3_5.forward(
- &input_ids,
- pixel_values,
- image_grid_thw,
- pixel_values_video,
- video_grid_thw,
- seqlen_offset,
- )?;
- let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
- 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..],
- )?
- };
- 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,
- Some(completion_tokens),
- Some(prompt_tokens),
+ let mut ctx = GenerationContext::new(
+ temperature.into(),
+ top_p.into(),
+ Some(20),
+ seed,
+ input_ids.dim(1)?,
+ sample_len,
+ self.device.clone(),
);
- Ok(response)
+ let data_vec = vec![
+ pixel_values,
+ image_grid_thw,
+ pixel_values_video,
+ video_grid_thw,
+ ];
+ let data = MultiModalData::new(data_vec);
+ generate_generic(
+ &mut self.qwen3_5,
+ &self.tokenizer,
+ input_ids,
+ data,
+ &mut ctx,
+ &self.model_name,
+ )
}
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 input = self.pre_processor.process_info(&mes, &mes_render)?;
+ let in_reasoning = mes_render.ends_with("\n");
let (mes_text, pixel_values, image_grid_thw, pixel_values_video, video_grid_thw) =
if let Some(processor) = &self.pre_processor {
let input = processor.process_info(&mes, &mes_render)?;
@@ -228,117 +195,30 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
} else {
(mes_render, None, None, None, None)
};
- let mut input_ids = self.tokenizer.text_encode(mes_text, &self.device)?;
- let mut seq_len = input_ids.dim(1)?;
- let mut seqlen_offset = 0;
+ let input_ids = self.tokenizer.text_encode(mes_text, &self.device)?;
let sample_len = mes.max_tokens.unwrap_or(1024);
- let stream = stream! {
- 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,
- image_grid_thw,
- pixel_values_video,
- video_grid_thw,
- seqlen_offset,
- )?;
- let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
- 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..],
- )?
- };
- 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_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;
- }
- "" => {
- // 结束工具调用
- 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
- );
- yield Ok(chunk);
- }
- }
- }
- 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();
- };
+ let data_vec = vec![
+ pixel_values,
+ image_grid_thw,
+ pixel_values_video,
+ video_grid_thw,
+ ];
+ let data = MultiModalData::new(data_vec);
+ let seed = mes.seed.unwrap_or(34562) as u64;
+ let stream = generate_stream_generic(
+ &mut self.qwen3_5,
+ &self.tokenizer,
+ input_ids,
+ data,
+ mes.temperature,
+ mes.top_p,
+ None,
+ seed,
+ sample_len,
+ in_reasoning,
+ &self.device,
+ &self.model_name,
+ )?;
Ok(Box::new(Box::pin(stream)))
}
}
diff --git a/src/models/qwen3_5/model.rs b/src/models/qwen3_5/model.rs
index 4d08a1a..7b0ffa9 100644
--- a/src/models/qwen3_5/model.rs
+++ b/src/models/qwen3_5/model.rs
@@ -10,6 +10,7 @@ use candle_nn::{
use crate::{
models::{
common::{
+ InferenceModel,
gguf::{GateUpDownMLPGguf, Gguf, ProjKind, QuantizedLinear},
modules::{conv1d_depthwise, eager_attention_forward, get_conv1d, softplus},
},
@@ -1043,10 +1044,11 @@ pub struct Qwen3_5Model {
language_model: Qwen3_5TextModel,
lm_head: ProjKind,
rope_deltas: Option,
+ stop_token_ids: Vec,
}
impl Qwen3_5Model {
- pub fn new_from_vb(vb: VarBuilder, config: Qwen3_5Config) -> Result {
+ pub fn new_from_vb(vb: VarBuilder, config: Qwen3_5Config, eos_ids: Vec) -> Result {
let vb_m = vb.pp("model");
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?;
let language_model =
@@ -1069,6 +1071,7 @@ impl Qwen3_5Model {
language_model,
lm_head: ProjKind::LinearProj(lm_head),
rope_deltas: None,
+ stop_token_ids: eos_ids,
})
}
@@ -1076,6 +1079,7 @@ impl Qwen3_5Model {
gguf: &mut Gguf,
mmproj_gguf: Option<&mut Gguf>,
device: &Device,
+ eos_ids: Vec,
) -> Result {
let spatial_merge_size = 2usize;
let image_token_id = 248056u32;
@@ -1102,6 +1106,7 @@ impl Qwen3_5Model {
language_model,
lm_head: ProjKind::QuantizedProj(QuantizedLinear::new(lm_head, None)),
rope_deltas: None,
+ stop_token_ids: eos_ids,
})
}
@@ -1434,6 +1439,46 @@ impl Qwen3_5Model {
}
pub fn clear_cache(&mut self) {
+ self.rope_deltas = None;
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 {
+ 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 {
+ self.forward(input_ids, None, None, None, None, seqlen_offset)
+ }
+
+ fn clear_cache(&mut self) {
+ self.clear_cache();
+ }
+
+ fn stop_token_ids(&self) -> Vec {
+ self.stop_token_ids.clone()
+ }
+}
diff --git a/src/models/qwen3_asr/config.rs b/src/models/qwen3_asr/config.rs
index 853f186..579124a 100644
--- a/src/models/qwen3_asr/config.rs
+++ b/src/models/qwen3_asr/config.rs
@@ -202,7 +202,7 @@ pub struct Qwen3ASRRopeScaling {
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
pub struct Qwen3ASRGenerationConfig {
pub do_sample: bool,
- pub eos_token_id: Vec,
+ pub eos_token_id: Vec,
pub pad_token_id: usize,
pub temperature: f32,
}
diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs
index d7a881e..019ec0a 100644
--- a/src/models/qwen3_asr/generate.rs
+++ b/src/models/qwen3_asr/generate.rs
@@ -1,6 +1,9 @@
+use std::time::Instant;
+
use crate::{
models::common::generate::get_logit_processor,
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
+ utils::response_utils::{build_chunk_response_with_usage, build_completion_response_with_time},
};
use anyhow::{Result, anyhow};
use candle_core::{DType, Device, Tensor};
@@ -21,8 +24,7 @@ use crate::{
},
tokenizer::TokenizerModel,
utils::{
- build_completion_chunk_response, build_completion_response, find_type_files, get_device,
- get_dtype,
+ find_type_files, get_device, get_dtype, response_utils::build_completion_chunk_response,
},
};
@@ -92,6 +94,8 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
let sample_len = mes.max_tokens.unwrap_or(1024);
let mut generate = Vec::new();
let mut prompt_tokens = 0u32;
+ let mut prompt_secs = 0.0f64;
+ let mut completion_secs = 0.0f64;
for data in audio_datas.iter() {
let mut input_ids = data.input_ids.clone();
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;
let mut seqlen_offset = 0;
for _ in 0..sample_len {
+ let i_start = Instant::now();
let logits =
self.qwen3_asr
.forward(&input_ids, seqlen_offset, input_features.as_ref())?;
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
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);
if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
break;
@@ -117,8 +128,14 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
}
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));
+ let response = build_completion_response_with_time(
+ res,
+ &self.model_name,
+ num_token.into(),
+ completion_secs.into(),
+ prompt_tokens.into(),
+ prompt_secs.into(),
+ );
Ok(response)
}
@@ -145,17 +162,30 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
let sample_len = mes.max_tokens.unwrap_or(1024);
let stream = stream! {
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() {
let mut input_ids = data.input_ids.clone();
let mut input_features = Some(data.input_features.clone().to_dtype(self.dtype)?);
let mut seq_len = input_ids.dim(1)?;
+ prompt_tokens += seq_len as u32;
let mut seqlen_offset = 0;
for _ in 0..sample_len {
+ let i_start = Instant::now();
let logits =
self.qwen3_asr
.forward(&input_ids, seqlen_offset, input_features.as_ref())?;
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
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();
if !error_tokens.is_empty() {
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);
yield Ok(chunk);
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;
}
seqlen_offset += seq_len;
diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs
index 16760cd..f4b91ca 100644
--- a/src/models/qwen3vl/generate.rs
+++ b/src/models/qwen3vl/generate.rs
@@ -1,11 +1,13 @@
use crate::{
- models::common::generate::get_logit_processor,
+ models::common::{
+ MultiModalData,
+ generate::{GenerationContext, generate_generic, generate_stream_generic},
+ },
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
};
-use anyhow::{Result, anyhow};
+use anyhow::Result;
use candle_core::{DType, Device, Tensor};
use candle_nn::VarBuilder;
-use rocket::async_stream::stream;
use rocket::futures::Stream;
use crate::{
@@ -16,10 +18,7 @@ use crate::{
qwen3vl::{config::Qwen3VLConfig, model::Qwen3VLModel, processor::Qwen3VLProcessor},
},
tokenizer::TokenizerModel,
- utils::{
- build_completion_chunk_response, build_completion_response, find_type_files, get_device,
- get_dtype,
- },
+ utils::{find_type_files, get_device, get_dtype},
};
pub struct Qwen3VLGenerateModel<'a> {
@@ -28,8 +27,6 @@ pub struct Qwen3VLGenerateModel<'a> {
pre_processor: Qwen3VLProcessor,
qwen3_vl: Qwen3VLModel,
device: Device,
- eos_token_id1: u32,
- eos_token_id2: u32,
generation_config: Qwen3GenerationConfig,
model_name: String,
}
@@ -46,10 +43,11 @@ impl<'a> Qwen3VLGenerateModel<'a> {
let pre_processor = Qwen3VLProcessor::new(path, &device, dtype)?;
let model_list = find_type_files(path, "safetensors")?;
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: Qwen3GenerationConfig =
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)
.file_name()
.and_then(|s| s.to_str())
@@ -61,8 +59,6 @@ impl<'a> Qwen3VLGenerateModel<'a> {
pre_processor,
qwen3_vl,
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,
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_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 enable_thinking = extract_metadata_value::(&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 mut input_ids = self
+ let input_ids = self
.tokenizer
.text_encode(input.replace_text.clone(), &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 = 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 seq_len = input_ids.dim(1)?;
let sample_len = mes.max_tokens.unwrap_or(1024);
- 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)?;
- generate.push(next_token);
- 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;
- }
- let num_token = generate.len() as u32;
- let res = self.tokenizer.token_decode(generate)?;
- self.qwen3_vl.clear_kv_cache();
- let response =
- build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
- Ok(response)
+ let mut ctx = GenerationContext::new(
+ temperature.into(),
+ top_p.into(),
+ top_k.into(),
+ seed,
+ input_ids.dim(1)?,
+ sample_len,
+ self.device.clone(),
+ );
+ let cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?;
+ let data_vec = vec![
+ input.pixel_values,
+ input.image_grid_thw,
+ input.pixel_values_video,
+ input.video_grid_thw,
+ cache_position.into(),
+ ];
+ let data = MultiModalData::new(data_vec);
+ generate_generic(
+ &mut self.qwen3_vl,
+ &self.tokenizer,
+ input_ids,
+ data,
+ &mut ctx,
+ &self.model_name,
+ )
}
fn generate_stream(
@@ -145,119 +124,141 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
.unwrap_or(self.generation_config.temperature);
let top_p = mes.top_p.unwrap_or(self.generation_config.top_p);
let top_k = self.generation_config.top_k;
- let seed = mes.seed.unwrap_or(34562) as u64;
- let 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 in_reasoning = mes_render.ends_with("\n");
let input = self.pre_processor.process_info(&mes, &mes_render)?;
- let mut input_ids = self
+ let input_ids = self
.tokenizer
.text_encode(input.replace_text.clone(), &self.device)?;
- let mut seq_len = input_ids.dim(1)?;
- let mut seqlen_offset = 0;
- let mut cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?;
+ let seq_len = input_ids.dim(1)?;
+ let cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?;
let sample_len = mes.max_tokens.unwrap_or(1024);
- let 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_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;
- }
- "" => {
- // 结束工具调用
- 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();
- };
+ let data_vec = vec![
+ input.pixel_values,
+ input.image_grid_thw,
+ input.pixel_values_video,
+ input.video_grid_thw,
+ cache_position.into(),
+ ];
+ let data = MultiModalData::new(data_vec);
+ let seed = mes.seed.unwrap_or(34562) as u64;
+ let stream = generate_stream_generic(
+ &mut self.qwen3_vl,
+ &self.tokenizer,
+ input_ids,
+ data,
+ temperature.into(),
+ top_p.into(),
+ top_k.into(),
+ seed,
+ sample_len,
+ in_reasoning,
+ &self.device,
+ &self.model_name,
+ )?;
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_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;
+ // }
+ // "" => {
+ // // 结束工具调用
+ // 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)))
}
}
diff --git a/src/models/qwen3vl/model.rs b/src/models/qwen3vl/model.rs
index c166d18..1e21a7f 100644
--- a/src/models/qwen3vl/model.rs
+++ b/src/models/qwen3vl/model.rs
@@ -10,6 +10,7 @@ use candle_nn::{
use crate::{
models::{
common::{
+ InferenceModel,
gguf::{Gguf, ProjKind, TwoLinearMLPGguf},
modules::{eager_attention_forward, get_layer_norm},
},
@@ -839,10 +840,11 @@ pub struct Qwen3VLModel {
language_model: Qwen3VLTextModel,
lm_head: Linear,
rope_deltas: Option,
+ stop_token_ids: Vec,
}
impl Qwen3VLModel {
- pub fn new(config: Qwen3VLConfig, vb: VarBuilder) -> Result {
+ pub fn new(config: Qwen3VLConfig, vb: VarBuilder, eos_ids: Vec) -> Result {
let vb_m = vb.pp("model");
let config = config.clone();
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?;
@@ -863,6 +865,7 @@ impl Qwen3VLModel {
language_model,
lm_head,
rope_deltas: None,
+ stop_token_ids: eos_ids,
})
}
@@ -1240,6 +1243,15 @@ impl Qwen3VLModel {
.broadcast_add(rope_deltas)?
.contiguous()?
.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 {
Tensor::zeros(1, inputs_embeds.dtype(), inputs_embeds.device())?
};
@@ -1265,6 +1277,48 @@ impl Qwen3VLModel {
}
pub fn clear_kv_cache(&mut self) {
+ self.rope_deltas = None;
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 {
+ 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 {
+ 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 {
+ self.stop_token_ids.clone()
+ }
+}
diff --git a/src/models/rmbg2_0/generate.rs b/src/models/rmbg2_0/generate.rs
index 28db495..a8f25e7 100644
--- a/src/models/rmbg2_0/generate.rs
+++ b/src/models/rmbg2_0/generate.rs
@@ -14,8 +14,9 @@ use rocket::futures::{Stream, stream};
use crate::{
models::{GenerateModel, rmbg2_0::model::BiRefNet},
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},
+ response_utils::build_img_completion_response,
},
};
diff --git a/src/models/voxcpm/generate.rs b/src/models/voxcpm/generate.rs
index 6dc4301..b12cb08 100644
--- a/src/models/voxcpm/generate.rs
+++ b/src/models/voxcpm/generate.rs
@@ -21,8 +21,8 @@ use crate::{
},
utils::{
audio_utils::{extract_audio_url, get_audio_wav_u8},
- build_audio_completion_response, extract_metadata_value, extract_user_text,
- find_type_files, get_device, get_dtype,
+ extract_metadata_value, extract_user_text, find_type_files, get_device, get_dtype,
+ response_utils::build_audio_completion_response,
},
};
diff --git a/src/params/chat.rs b/src/params/chat.rs
index 8747bc0..ea314d4 100644
--- a/src/params/chat.rs
+++ b/src/params/chat.rs
@@ -70,6 +70,8 @@ pub struct ChatCompletionParameters {
/// Developer-defined tags and values used for filtering completions in the dashboard.
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata: Option>,
+ #[serde(skip_serializing_if = "Option::is_none")]
+ pub enable_thinking: Option,
/// 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.
#[serde(skip_serializing_if = "Option::is_none")]
diff --git a/src/utils/mod.rs b/src/utils/mod.rs
index b930531..ea1dc2f 100644
--- a/src/utils/mod.rs
+++ b/src/utils/mod.rs
@@ -1,6 +1,7 @@
pub mod audio_utils;
pub mod img_utils;
pub mod interpolate;
+pub mod response_utils;
pub mod tensor_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 crate::models::common::model_mapping::WhichModel;
-use crate::params::{
- chat::{
- AudioUrlType, ChatCompletionChoice, ChatCompletionChunkChoice, ChatCompletionChunkResponse,
- ChatCompletionParameters, ChatCompletionResponse, ChatMessage, ChatMessageAudioContentPart,
- ChatMessageContent, ChatMessageContentPart, ChatMessageImageContentPart, DeltaChatMessage,
- DeltaFunction, DeltaToolCall, Function, ImageUrlType, ToolCall,
- },
- shared::{FinishReason, Usage},
+use crate::params::chat::{
+ ChatCompletionParameters, ChatMessage, ChatMessageContent, ChatMessageContentPart,
};
use anyhow::{Result, anyhow};
use byteorder::{LittleEndian, ReadBytesExt};
@@ -409,305 +404,6 @@ pub fn ceil_by_factor(num: f32, factor: u32) -> u32 {
ceil * factor
}
-pub fn build_img_completion_response(
- base64vec: &Vec,
- 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) -> 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("") {
- let mes: Vec<&str> = res.split("").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("", "");
- let function = match serde_json::from_str::(&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,
- prompt_tokens: Option,
-) -> 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,
- completion_secs: Option,
- prompt_tokens: Option,
- prompt_secs: Option,
-) -> 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,
- tool_call_content: Option,
-) -> 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::(&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> {
let mut mes_vec = Vec::new();
for chat_mes in mes.messages.clone() {
diff --git a/src/utils/response_utils.rs b/src/utils/response_utils.rs
new file mode 100644
index 0000000..9ff24a2
--- /dev/null
+++ b/src/utils/response_utils.rs
@@ -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,
+ 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 "" where the content before the first
+/// delimiter is treated as the main content, and subsequent parts (separated by "")
+/// are parsed as JSON tool call definitions.
+/// 2. Reasoning content formatting where content between and tags is
+/// extracted as reasoning_content, and the content after 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 "" and then JSON-formatted tool call data separated by ""
+/// - Reasoning format: Content wrapped in and tags followed by actual response
+fn build_response(res: String, model_name: &str, usage: Option) -> 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("") {
+ let mes: Vec<&str> = res.split("").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("", "");
+ let function = match serde_json::from_str::(&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("") {
+ let contents: Vec<&str> = content.split("").collect();
+ let reasoning_content = contents[0].to_string().replace("", "");
+ 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,
+ prompt_tokens: Option,
+) -> 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,
+ completion_secs: Option,
+ prompt_tokens: Option,
+ prompt_secs: Option,
+) -> 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,
+ completion_secs: Option,
+ prompt_tokens: Option,
+ prompt_secs: Option,
+) -> 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,
+ tool_call_content: Option,
+) -> 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::(&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
+}
diff --git a/tests/test_deepseek_ocr.rs b/tests/test_deepseek_ocr.rs
index 0da515a..492d9c1 100644
--- a/tests/test_deepseek_ocr.rs
+++ b/tests/test_deepseek_ocr.rs
@@ -89,17 +89,11 @@ fn deepseek_ocr_generate() -> Result<()> {
let mut model = DeepseekOCRGenerateModel::init(&model_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
let res = model.generate(mes)?;
- let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
if let Some(usage) = &res.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!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
diff --git a/tests/test_fun_asr_nano.rs b/tests/test_fun_asr_nano.rs
index c36f222..c17cead 100644
--- a/tests/test_fun_asr_nano.rs
+++ b/tests/test_fun_asr_nano.rs
@@ -38,17 +38,11 @@ fn fun_asr_nano_generate() -> Result<()> {
let mut fun_asr_model = FunAsrNanoGenerateModel::init(&model_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
let res = fun_asr_model.generate(mes)?;
- let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
if let Some(usage) = &res.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!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
diff --git a/tests/test_gelab_zero.rs b/tests/test_gelab_zero.rs
index 131e227..6c61a02 100644
--- a/tests/test_gelab_zero.rs
+++ b/tests/test_gelab_zero.rs
@@ -62,16 +62,10 @@ fn gelab_zero_generate() -> Result<()> {
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
let res = qwen3vl.generate(mes)?;
- let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
if let Some(usage) = &res.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!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
diff --git a/tests/test_gguf_qwen3_5.rs b/tests/test_gguf_qwen3_5.rs
index c8133fc..f17f3a9 100644
--- a/tests/test_gguf_qwen3_5.rs
+++ b/tests/test_gguf_qwen3_5.rs
@@ -80,16 +80,10 @@ fn gguf_test() -> Result<()> {
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
let res = gguf_qwen3_5.generate(mes)?;
- let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
if let Some(usage) = &res.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!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
diff --git a/tests/test_glm_asr_nano.rs b/tests/test_glm_asr_nano.rs
index d0df313..0d75e64 100644
--- a/tests/test_glm_asr_nano.rs
+++ b/tests/test_glm_asr_nano.rs
@@ -39,23 +39,17 @@ fn glm_asr_nano_generate() -> Result<()> {
let mut glm_asr_model = GlmAsrNanoGenerateModel::init(&model_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
let res = glm_asr_model.generate(mes)?;
- let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
if let Some(usage) = &res.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!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
#[tokio::test]
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 =
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);
diff --git a/tests/test_glm_ocr.rs b/tests/test_glm_ocr.rs
index b9b7a6c..c425ba6 100644
--- a/tests/test_glm_ocr.rs
+++ b/tests/test_glm_ocr.rs
@@ -39,17 +39,11 @@ fn glm_ocr_generate() -> Result<()> {
let mut model = GlmOcrGenerateModel::init(&model_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
let res = model.generate(mes)?;
- let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
if let Some(usage) = &res.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!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
diff --git a/tests/test_health_models.rs b/tests/test_health_models.rs
deleted file mode 100644
index 8c3d569..0000000
--- a/tests/test_health_models.rs
+++ /dev/null
@@ -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");
-// }
diff --git a/tests/test_hunyuan_ocr.rs b/tests/test_hunyuan_ocr.rs
index c750cc7..48242e1 100644
--- a/tests/test_hunyuan_ocr.rs
+++ b/tests/test_hunyuan_ocr.rs
@@ -39,18 +39,11 @@ fn hunyuan_ocr_generate() -> Result<()> {
let mut model = HunyuanOCRGenerateModel::init(&model_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
let res = model.generate(mes)?;
- let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
if let Some(usage) = &res.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!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
-
Ok(())
}
diff --git a/tests/test_lfm2.rs b/tests/test_lfm2.rs
index 3c678a8..2659a35 100644
--- a/tests/test_lfm2.rs
+++ b/tests/test_lfm2.rs
@@ -31,17 +31,11 @@ fn lfm2_generate() -> Result<()> {
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
- let result = model.generate(mes)?;
- let i_duration = i_start.elapsed();
- println!("generate: \n {:?}", result);
- 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);
+ let res = model.generate(mes)?;
+ println!("generate: \n {:?}", res);
+ if let Some(usage) = &res.usage {
+ println!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
diff --git a/tests/test_lfm2vl.rs b/tests/test_lfm2vl.rs
index b095471..8ab7150 100644
--- a/tests/test_lfm2vl.rs
+++ b/tests/test_lfm2vl.rs
@@ -43,18 +43,11 @@ fn lfm2vl_generate() -> Result<()> {
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
- let result = model.generate(mes)?;
- let i_duration = i_start.elapsed();
- println!("generate: \n {:?}", result);
- 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);
+ let res = model.generate(mes)?;
+ println!("generate: \n {:?}", res);
+ if let Some(usage) = &res.usage {
+ println!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
-
Ok(())
}
diff --git a/tests/test_minicpm4.rs b/tests/test_minicpm4.rs
index 9e3753a..671a383 100644
--- a/tests/test_minicpm4.rs
+++ b/tests/test_minicpm4.rs
@@ -33,18 +33,11 @@ fn minicpm_generate() -> Result<()> {
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
- let result = model.generate(mes)?;
- let i_duration = i_start.elapsed();
- println!("generate: \n {:?}", result);
- 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);
+ let res = model.generate(mes)?;
+ println!("generate: \n {:?}", res);
+ if let Some(usage) = &res.usage {
+ println!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
-
Ok(())
}
diff --git a/tests/test_paddleocr_vl.rs b/tests/test_paddleocr_vl.rs
index 640ad5a..a6ff2e6 100644
--- a/tests/test_paddleocr_vl.rs
+++ b/tests/test_paddleocr_vl.rs
@@ -89,17 +89,11 @@ fn paddleocr_vl_generate() -> Result<()> {
let mut model = PaddleOCRVLGenerateModel::init(&model_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
let res = model.generate(mes)?;
- let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
if let Some(usage) = &res.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!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
diff --git a/tests/test_qwen2_5vl.rs b/tests/test_qwen2_5vl.rs
index b4ae2f6..120edd8 100644
--- a/tests/test_qwen2_5vl.rs
+++ b/tests/test_qwen2_5vl.rs
@@ -46,18 +46,11 @@ fn qwen2_5vl_generate() -> Result<()> {
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
- let result = model.generate(mes)?;
- let i_duration = i_start.elapsed();
- println!("generate: \n {:?}", result);
- 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);
+ let res = model.generate(mes)?;
+ println!("generate: \n {:?}", res);
+ if let Some(usage) = &res.usage {
+ println!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
-
Ok(())
}
diff --git a/tests/test_qwen3.rs b/tests/test_qwen3.rs
index fd94b51..895daf0 100644
--- a/tests/test_qwen3.rs
+++ b/tests/test_qwen3.rs
@@ -21,7 +21,8 @@ fn qwen3_0_6b_generate() -> Result<()> {
"role": "user",
"content": "你好啊,你是谁"
}
- ]
+ ],
+ "enable_thinking": true
}
"#;
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
@@ -30,17 +31,11 @@ fn qwen3_0_6b_generate() -> Result<()> {
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
- let result = model.generate(mes)?;
- let i_duration = i_start.elapsed();
- println!("generate: \n {:?}", result);
- 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);
+ let res = model.generate(mes)?;
+ println!("generate: \n {:?}", res);
+ if let Some(usage) = &res.usage {
+ println!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
@@ -61,7 +56,8 @@ async fn qwen3_0_6b_stream() -> Result<()> {
"role": "user",
"content": "你是谁"
}
- ]
+ ],
+ "enable_thinking": true
}
"#;
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
diff --git a/tests/test_qwen3_5.rs b/tests/test_qwen3_5.rs
index d901890..176192a 100644
--- a/tests/test_qwen3_5.rs
+++ b/tests/test_qwen3_5.rs
@@ -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 i_start = Instant::now();
let mut qwen3_5 = Qwen3_5GenerateModel::init(&model_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
let res = qwen3_5.generate(mes)?;
- let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
if let Some(usage) = &res.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!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
@@ -75,16 +71,17 @@ async fn qwen3_5_stream() -> Result<()> {
"type": "image",
"image_url":
{
- "url": "file:///home/jhq/Downloads/gougou1.jpg"
+ "url": "file://./assets/img/ocr_test3.png"
}
},
{
"type": "text",
- "text": "描述这张图片."
+ "text": "OCR"
}
]
}
- ]
+ ],
+ "enable_thinking": true
}
"#;
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
diff --git a/tests/test_qwen3_asr.rs b/tests/test_qwen3_asr.rs
index f477b64..5e750c6 100644
--- a/tests/test_qwen3_asr.rs
+++ b/tests/test_qwen3_asr.rs
@@ -35,17 +35,11 @@ fn qwen3_asr_generate() -> Result<()> {
let mut model = Qwen3AsrGenerateModel::init(&model_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
let res = model.generate(mes)?;
- let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
if let Some(usage) = &res.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!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
diff --git a/tests/test_qwen3vl.rs b/tests/test_qwen3vl.rs
index 039096a..7d0fd61 100644
--- a/tests/test_qwen3vl.rs
+++ b/tests/test_qwen3vl.rs
@@ -94,17 +94,11 @@ fn qwen3vl_generate() -> Result<()> {
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
let res = qwen3vl.generate(mes)?;
- let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
if let Some(usage) = &res.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!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
diff --git a/tests/test_robo_brain.rs b/tests/test_robo_brain.rs
index 693e2d2..7da265f 100644
--- a/tests/test_robo_brain.rs
+++ b/tests/test_robo_brain.rs
@@ -34,17 +34,10 @@ fn robo_brain_generate() -> Result<()> {
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
- let i_start = Instant::now();
- let result = model.generate(mes)?;
- let i_duration = i_start.elapsed();
- println!("generate: \n {:?}", result);
- 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);
+ let res = model.generate(mes)?;
+ println!("generate: \n {:?}", res);
+ if let Some(usage) = &res.usage {
+ println!("usage: \n {:?}", usage);
}
- println!("Time elapsed in generate is: {:?}", i_duration);
-
Ok(())
}