diff --git a/src/models/deepseek_ocr/generate.rs b/src/models/deepseek_ocr/generate.rs
index fc14c4e..3904758 100644
--- a/src/models/deepseek_ocr/generate.rs
+++ b/src/models/deepseek_ocr/generate.rs
@@ -127,9 +127,10 @@ impl GenerateModel for DeepseekOCRGenerateModel {
images_seq_mask = None;
images_spatial_crop_t = None;
}
+ let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?;
self.deepseekocr_model.clear_kv_cache();
- let response = build_completion_response(res, &self.model_name);
+ let response = build_completion_response(res, &self.model_name, Some(num_token));
Ok(response)
}
diff --git a/src/models/hunyuan_ocr/generate.rs b/src/models/hunyuan_ocr/generate.rs
index 0f35df4..04295d9 100644
--- a/src/models/hunyuan_ocr/generate.rs
+++ b/src/models/hunyuan_ocr/generate.rs
@@ -119,9 +119,10 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
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);
+ let response = build_completion_response(res, &self.model_name, Some(num_token));
Ok(response)
}
diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs
index b8af4c3..100cdcf 100644
--- a/src/models/minicpm4/generate.rs
+++ b/src/models/minicpm4/generate.rs
@@ -78,9 +78,10 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
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);
+ let response = build_completion_response(res, &self.model_name, Some(num_token));
Ok(response)
}
fn generate_stream(
diff --git a/src/models/paddleocr_vl/generate.rs b/src/models/paddleocr_vl/generate.rs
index 254f377..36e470d 100644
--- a/src/models/paddleocr_vl/generate.rs
+++ b/src/models/paddleocr_vl/generate.rs
@@ -104,9 +104,10 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
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);
+ let response = build_completion_response(res, &self.model_name, Some(num_token));
Ok(response)
}
diff --git a/src/models/qwen2_5vl/generate.rs b/src/models/qwen2_5vl/generate.rs
index 08b41ba..7dd65a7 100644
--- a/src/models/qwen2_5vl/generate.rs
+++ b/src/models/qwen2_5vl/generate.rs
@@ -118,9 +118,10 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
pixel_values = None;
pixel_values_video = None;
}
+ 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);
+ let response = build_completion_response(res, &self.model_name, Some(num_token));
Ok(response)
}
@@ -167,6 +168,8 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
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();
for _ in 0..sample_len {
let logits = self.qwen2_5_vl.forward(
&input_ids,
@@ -203,8 +206,56 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
continue;
}
error_tokens.clear();
- let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
- yield Ok(chunk);
+ // 处理特殊标记和工具调用
+ 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);
+ }
+ }
+ }
+ // 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;
}
diff --git a/src/models/qwen2_5vl/processor.rs b/src/models/qwen2_5vl/processor.rs
index e15b696..58d8039 100644
--- a/src/models/qwen2_5vl/processor.rs
+++ b/src/models/qwen2_5vl/processor.rs
@@ -72,6 +72,9 @@ impl Qwen2_5VLProcessor {
if let ChatMessageContentPart::Image(img_part) = part {
let img_url = img_part.image_url;
vision_map.get_mut("image").unwrap().push(img_url.url);
+ // } else if let ChatMessageContentPart::Video(video_part) = part {
+ // let video_url = video_part.video_url;
+ // vision_map.get_mut("video").unwrap().push(video_url.url);
}
}
}
diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs
index a5c3de1..25185b8 100644
--- a/src/models/qwen3vl/generate.rs
+++ b/src/models/qwen3vl/generate.rs
@@ -120,9 +120,10 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
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);
+ let response = build_completion_response(res, &self.model_name, Some(num_token));
Ok(response)
}
@@ -174,18 +175,18 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
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();
+ 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);
}
@@ -206,27 +207,54 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
}
error_tokens.clear();
- if decoded_token.as_str() == "" {
- tool_call_id = Some(uuid::Uuid::new_v4().to_string());
- continue;
- } else {
- if decoded_token.as_str() == "" {
- let chunk = build_completion_chunk_response(decoded_token, &self.model_name, tool_call_id.clone(), Some(tool_call_content.clone()));
+ // 处理特殊标记和工具调用
+ 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);
}
- else {
+ _ => {
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);
+ } 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;
}
diff --git a/src/models/qwen3vl/model.rs b/src/models/qwen3vl/model.rs
index ca65598..a84610b 100644
--- a/src/models/qwen3vl/model.rs
+++ b/src/models/qwen3vl/model.rs
@@ -730,7 +730,6 @@ impl Qwen3VLTextModel {
deepstack_visual_embeds: Option>,
) -> Result {
let (b_size, seq_len, _) = inputs_embeds.dims3()?;
-
let position_ids = match position_ids {
Some(ids) => ids.clone(),
None => Tensor::arange(
diff --git a/src/utils/mod.rs b/src/utils/mod.rs
index 653689f..b90d19b 100644
--- a/src/utils/mod.rs
+++ b/src/utils/mod.rs
@@ -11,7 +11,7 @@ use aha_openai_dive::v1::resources::{
ChatCompletionParameters, ChatCompletionResponse, ChatMessage, ChatMessageContent,
ChatMessageContentPart, DeltaChatMessage, DeltaFunction, DeltaToolCall, Function, ToolCall,
},
- shared::FinishReason,
+ shared::{FinishReason, Usage},
};
use anyhow::Result;
use candle_core::{DType, Device};
@@ -137,8 +137,24 @@ pub fn ceil_by_factor(num: f32, factor: u32) -> u32 {
ceil * factor
}
-pub fn build_completion_response(res: String, model_name: &str) -> ChatCompletionResponse {
+pub fn build_completion_response(res: String, model_name: &str, num_tokens: Option) -> ChatCompletionResponse {
let id = uuid::Uuid::new_v4().to_string();
+ let usage = match num_tokens {
+ Some(num) => {
+ Some(Usage {
+ input_tokens: None,
+ input_tokens_details: None,
+ output_tokens: None,
+ output_tokens_details: None,
+ prompt_tokens: None,
+ completion_tokens: None,
+ total_tokens: num,
+ prompt_tokens_details: None,
+ completion_tokens_details: None
+ })
+ },
+ None => None,
+ };
let mut response = ChatCompletionResponse {
id: Some(id),
choices: vec![],
@@ -147,7 +163,7 @@ pub fn build_completion_response(res: String, model_name: &str) -> ChatCompletio
service_tier: None,
system_fingerprint: None,
object: "chat.completion".to_string(),
- usage: None,
+ usage
};
let choice = if res.contains("") {
let mes: Vec<&str> = res.split("").collect();
diff --git a/tests/test_deepseek_ocr.rs b/tests/test_deepseek_ocr.rs
index 115fc07..6274993 100644
--- a/tests/test_deepseek_ocr.rs
+++ b/tests/test_deepseek_ocr.rs
@@ -41,8 +41,14 @@ fn deepseek_ocr_generate() -> Result<()> {
let i_start = Instant::now();
let res = model.generate(mes)?;
let i_duration = i_start.elapsed();
- println!("Time elapsed in generate is: {:?}", i_duration);
println!("generate: \n {:?}", res);
+ if res.usage.is_some() {
+ let num_token = res.usage.as_ref().unwrap().total_tokens;
+ let duration_secs = i_duration.as_secs_f64();
+ let tps = num_token as f64 / duration_secs;
+ println!("Tokens per second (TPS): {:.2}", tps);
+ }
+ println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
diff --git a/tests/test_gelab_zero.rs b/tests/test_gelab_zero.rs
index 06cfe3f..4aff3e8 100644
--- a/tests/test_gelab_zero.rs
+++ b/tests/test_gelab_zero.rs
@@ -64,6 +64,12 @@ fn gelab_zero_generate() -> Result<()> {
let res = qwen3vl.generate(mes)?;
let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
+ if res.usage.is_some() {
+ let num_token = res.usage.as_ref().unwrap().total_tokens;
+ let duration_secs = i_duration.as_secs_f64();
+ let tps = num_token as f64 / duration_secs;
+ println!("Tokens per second (TPS): {:.2}", tps);
+ }
println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
diff --git a/tests/test_hunyuan_ocr.rs b/tests/test_hunyuan_ocr.rs
index 5e8c9e1..5b14ba8 100644
--- a/tests/test_hunyuan_ocr.rs
+++ b/tests/test_hunyuan_ocr.rs
@@ -40,8 +40,15 @@ fn hunyuan_ocr_generate() -> Result<()> {
let i_start = Instant::now();
let res = model.generate(mes)?;
let i_duration = i_start.elapsed();
- println!("Time elapsed in generate is: {:?}", i_duration);
println!("generate: \n {:?}", res);
+ if res.usage.is_some() {
+ let num_token = res.usage.as_ref().unwrap().total_tokens;
+ let duration_secs = i_duration.as_secs_f64();
+ let tps = num_token as f64 / duration_secs;
+ println!("Tokens per second (TPS): {:.2}", tps);
+ }
+ println!("Time elapsed in generate is: {:?}", i_duration);
+
Ok(())
}
diff --git a/tests/test_minicpm4.rs b/tests/test_minicpm4.rs
index 0d98fd2..5477e17 100644
--- a/tests/test_minicpm4.rs
+++ b/tests/test_minicpm4.rs
@@ -33,8 +33,14 @@ fn minicpm_generate() -> Result<()> {
let i_start = Instant::now();
let result = model.generate(mes)?;
- println!("generate: \n {:?}", result);
let i_duration = i_start.elapsed();
+ println!("generate: \n {:?}", result);
+ if result.usage.is_some() {
+ let num_token = result.usage.as_ref().unwrap().total_tokens;
+ let duration_secs = i_duration.as_secs_f64();
+ let tps = num_token as f64 / duration_secs;
+ println!("Tokens per second (TPS): {:.2}", tps);
+ }
println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
diff --git a/tests/test_paddleocr_vl.rs b/tests/test_paddleocr_vl.rs
index 738f686..571cbfd 100644
--- a/tests/test_paddleocr_vl.rs
+++ b/tests/test_paddleocr_vl.rs
@@ -41,8 +41,14 @@ fn paddleocr_vl_generate() -> Result<()> {
let i_start = Instant::now();
let res = model.generate(mes)?;
let i_duration = i_start.elapsed();
- println!("Time elapsed in generate is: {:?}", i_duration);
println!("generate: \n {:?}", res);
+ if res.usage.is_some() {
+ let num_token = res.usage.as_ref().unwrap().total_tokens;
+ let duration_secs = i_duration.as_secs_f64();
+ let tps = num_token as f64 / duration_secs;
+ println!("Tokens per second (TPS): {:.2}", tps);
+ }
+ println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
diff --git a/tests/test_qwen2_5vl.rs b/tests/test_qwen2_5vl.rs
index 1159e95..18a6ead 100644
--- a/tests/test_qwen2_5vl.rs
+++ b/tests/test_qwen2_5vl.rs
@@ -45,9 +45,15 @@ fn qwen2_5vl_generate() -> Result<()> {
println!("Time elapsed in load model is: {:?}", i_duration);
let i_start = Instant::now();
- let result = model.generate(mes)?;
- println!("generate: \n {:?}", result);
+ let result = model.generate(mes)?;
let i_duration = i_start.elapsed();
+ println!("generate: \n {:?}", result);
+ if result.usage.is_some() {
+ let num_token = result.usage.as_ref().unwrap().total_tokens;
+ let duration_secs = i_duration.as_secs_f64();
+ let tps = num_token as f64 / duration_secs;
+ println!("Tokens per second (TPS): {:.2}", tps);
+ }
println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
diff --git a/tests/test_qwen3vl.rs b/tests/test_qwen3vl.rs
index 12eb79d..f7a5ce7 100644
--- a/tests/test_qwen3vl.rs
+++ b/tests/test_qwen3vl.rs
@@ -19,15 +19,15 @@ fn qwen3vl_generate() -> Result<()> {
"role": "user",
"content": [
{
- "type": "image",
- "image_url":
+ "type": "video",
+ "video_url":
{
- "url": "file://./assets/img/ocr_test1.png"
+ "url": "./assets/video/video_test.mp4"
}
- },
+ },
{
"type": "text",
- "text": "请分析图片并提取所有可见文本内容,按从左到右、从上到下的布局,返回纯文本"
+ "text": "视频中发生了什么?, 现在几点了"
}
]
}
@@ -44,13 +44,19 @@ fn qwen3vl_generate() -> Result<()> {
let res = qwen3vl.generate(mes)?;
let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
+ if res.usage.is_some() {
+ let num_token = res.usage.as_ref().unwrap().total_tokens;
+ let duration_secs = i_duration.as_secs_f64();
+ let tps = num_token as f64 / duration_secs;
+ println!("Tokens per second (TPS): {:.2}", tps);
+ }
println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
}
#[tokio::test]
async fn qwen3vl_stream() -> Result<()> {
- // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda qwen3vl_stream -r -- --nocapture
+ // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda,ffmpeg qwen3vl_stream -r -- --nocapture
let model_path = "/home/jhq/huggingface_model/Qwen/Qwen3-VL-2B-Instruct/";
@@ -60,7 +66,7 @@ async fn qwen3vl_stream() -> Result<()> {
"messages": [
{
"role": "user",
- "content": [
+ "content": [
{
"type": "video",
"video_url":
@@ -70,39 +76,11 @@ async fn qwen3vl_stream() -> Result<()> {
},
{
"type": "text",
- "text": "视频中发生了什么?, 现在几点了"
+ "text": "视频中发生了什么?"
}
]
}
- ],
- "tools": [
- {
- "type": "function",
- "function": {
- "name": "get_current_time",
- "description": "当你想知道现在的时间时非常有用。",
- "parameters": {}
- }
- },
- {
- "type": "function",
- "function": {
- "name": "get_current_weather",
- "description": "当你想查询指定城市的天气时非常有用。",
- "parameters": {
- "type": "object",
- "properties": {
- "location": {
- "type": "string",
- "description": "城市或县区,比如北京市、杭州市、余杭区等。"
- }
- },
- "required": ["location"]
- }
- }
- }
- ],
- "tool_choice": null
+ ]
}
"#;
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
diff --git a/tests/test_robo_brain.rs b/tests/test_robo_brain.rs
index dd27672..2cd02a7 100644
--- a/tests/test_robo_brain.rs
+++ b/tests/test_robo_brain.rs
@@ -34,8 +34,14 @@ fn robo_brain_generate() -> Result<()> {
let i_start = Instant::now();
let result = model.generate(mes)?;
- println!("generate: \n {:?}", result);
let i_duration = i_start.elapsed();
+ println!("generate: \n {:?}", result);
+ if result.usage.is_some() {
+ let num_token = result.usage.as_ref().unwrap().total_tokens;
+ let duration_secs = i_duration.as_secs_f64();
+ let tps = num_token as f64 / duration_secs;
+ println!("Tokens per second (TPS): {:.2}", tps);
+ }
println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())