From 3eb983b5bbc2a69880461d9108b164eb901e5d1b Mon Sep 17 00:00:00 2001 From: jhqxxx Date: Wed, 10 Dec 2025 12:57:24 +0800 Subject: [PATCH 1/2] stash save --- src/chat_template/mod.rs | 1 + src/models/qwen3vl/generate.rs | 26 ++++++++++++++++++++-- tests/test_gelab_zero.rs | 32 +++++++++++++++++++++++++-- tests/test_qwen3vl.rs | 40 +++++++++++++++++++++++++++++----- 4 files changed, 89 insertions(+), 10 deletions(-) diff --git a/src/chat_template/mod.rs b/src/chat_template/mod.rs index 24c4bf7..48aefb4 100644 --- a/src/chat_template/mod.rs +++ b/src/chat_template/mod.rs @@ -102,6 +102,7 @@ impl<'a> ChatTemplate<'a> { pub fn apply_chat_template(&self, messages: &ChatCompletionParameters) -> Result { let context = context! { messages => &messages.messages, + tools => &messages.tools.as_ref(), add_generation_prompt => true, }; let template = self diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs index 775b822..a5c3de1 100644 --- a/src/models/qwen3vl/generate.rs +++ b/src/models/qwen3vl/generate.rs @@ -171,6 +171,8 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'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.qwen3_vl.forward( &input_ids, @@ -203,8 +205,28 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { continue; } error_tokens.clear(); - let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None); - yield Ok(chunk); + + 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())); + 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); + 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; } diff --git a/tests/test_gelab_zero.rs b/tests/test_gelab_zero.rs index 29db2d7..06cfe3f 100644 --- a/tests/test_gelab_zero.rs +++ b/tests/test_gelab_zero.rs @@ -19,11 +19,39 @@ fn gelab_zero_generate() -> Result<()> { "content": [ { "type": "text", - "text": "Hello, GELab-Zero!" + "text": "Hello, GELab-Zero!, 现在几点了" } ] } - ] + ], + "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_qwen3vl.rs b/tests/test_qwen3vl.rs index 609ee59..12eb79d 100644 --- a/tests/test_qwen3vl.rs +++ b/tests/test_qwen3vl.rs @@ -19,15 +19,15 @@ fn qwen3vl_generate() -> Result<()> { "role": "user", "content": [ { - "type": "video", - "video_url": + "type": "image", + "image_url": { - "url": "./assets/video/video_test.mp4" + "url": "file://./assets/img/ocr_test1.png" } }, { "type": "text", - "text": "视频里发生了什么" + "text": "请分析图片并提取所有可见文本内容,按从左到右、从上到下的布局,返回纯文本" } ] } @@ -70,11 +70,39 @@ 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)?; From 1a0df13bc7fed21badaabb1ccd0cbf136ff2329c Mon Sep 17 00:00:00 2001 From: jhqxxx Date: Wed, 10 Dec 2025 16:44:20 +0800 Subject: [PATCH 2/2] model test add tps --- src/models/deepseek_ocr/generate.rs | 3 +- src/models/hunyuan_ocr/generate.rs | 3 +- src/models/minicpm4/generate.rs | 3 +- src/models/paddleocr_vl/generate.rs | 3 +- src/models/qwen2_5vl/generate.rs | 57 ++++++++++++++++++++-- src/models/qwen2_5vl/processor.rs | 3 ++ src/models/qwen3vl/generate.rs | 76 ++++++++++++++++++++--------- src/models/qwen3vl/model.rs | 1 - src/utils/mod.rs | 22 +++++++-- tests/test_deepseek_ocr.rs | 8 ++- tests/test_gelab_zero.rs | 6 +++ tests/test_hunyuan_ocr.rs | 9 +++- tests/test_minicpm4.rs | 8 ++- tests/test_paddleocr_vl.rs | 8 ++- tests/test_qwen2_5vl.rs | 10 +++- tests/test_qwen3vl.rs | 52 ++++++-------------- tests/test_robo_brain.rs | 8 ++- 17 files changed, 201 insertions(+), 79 deletions(-) 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(())