usage add prompt_tokens

This commit is contained in:
jhqxxx
2026-03-06 18:38:16 +08:00
parent 2d37efaab6
commit 548eb8185e
12 changed files with 77 additions and 28 deletions
+3 -1
View File
@@ -102,6 +102,7 @@ impl GenerateModel for DeepseekOCRGenerateModel {
let mut images_spatial_crop_t = Some(&images_spatial_crop_t); let mut images_spatial_crop_t = Some(&images_spatial_crop_t);
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut generate = Vec::new(); let mut generate = Vec::new();
let sample_len = mes.max_tokens.unwrap_or(1024); let sample_len = mes.max_tokens.unwrap_or(1024);
for _ in 0..sample_len { for _ in 0..sample_len {
@@ -130,7 +131,8 @@ impl GenerateModel for DeepseekOCRGenerateModel {
let num_token = generate.len() as u32; let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?; let res = self.tokenizer.token_decode(generate)?;
self.deepseekocr_model.clear_kv_cache(); self.deepseekocr_model.clear_kv_cache();
let response = build_completion_response(res, &self.model_name, Some(num_token)); let response =
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
Ok(response) Ok(response)
} }
+3 -1
View File
@@ -114,6 +114,7 @@ impl GenerateModel for FunAsrNanoGenerateModel {
let mut speech = Some(speech.to_dtype(self.dtype)?); let mut speech = Some(speech.to_dtype(self.dtype)?);
let mut fbank_mask = Some(&fbank_mask); let mut fbank_mask = Some(&fbank_mask);
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let mut generate = Vec::new(); let mut generate = Vec::new();
let sample_len = mes.max_tokens.unwrap_or(1024); let sample_len = mes.max_tokens.unwrap_or(1024);
@@ -139,7 +140,8 @@ impl GenerateModel for FunAsrNanoGenerateModel {
let num_token = generate.len() as u32; let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?; let res = self.tokenizer.token_decode(generate)?;
self.fun_asr_nano.clear_kv_cache(); self.fun_asr_nano.clear_kv_cache();
let response = build_completion_response(res, &self.model_name, Some(num_token)); let response =
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
Ok(response) Ok(response)
} }
+3 -1
View File
@@ -77,6 +77,7 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
let mut input_features = Some(input_features.to_dtype(self.dtype)?); let mut input_features = Some(input_features.to_dtype(self.dtype)?);
let mut audio_token_lengths = Some(audio_token_lengths); let mut audio_token_lengths = Some(audio_token_lengths);
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let mut generate: Vec<u32> = Vec::new(); let mut generate: Vec<u32> = Vec::new();
let sample_len = mes.max_tokens.unwrap_or(1024); let sample_len = mes.max_tokens.unwrap_or(1024);
@@ -105,7 +106,8 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
let num_token = generate.len() as u32; let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?; let res = self.tokenizer.token_decode(generate)?;
self.glm_asr_nano.clear_kv_cache(); self.glm_asr_nano.clear_kv_cache();
let response = build_completion_response(res, &self.model_name, Some(num_token)); let response =
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
Ok(response) Ok(response)
} }
+3 -1
View File
@@ -93,6 +93,7 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
let mut pixel_values = data.pixel_values; let mut pixel_values = data.pixel_values;
let mut image_grid_thw = data.image_grid_thw; let mut image_grid_thw = data.image_grid_thw;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let mut generate: Vec<u32> = Vec::new(); let mut generate: Vec<u32> = Vec::new();
let sample_len = mes.max_tokens.unwrap_or(1024); let sample_len = mes.max_tokens.unwrap_or(1024);
@@ -122,7 +123,8 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
let num_token = generate.len() as u32; let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?; let res = self.tokenizer.token_decode(generate)?;
self.hunyuan_vl.clear_kv_cache(); self.hunyuan_vl.clear_kv_cache();
let response = build_completion_response(res, &self.model_name, Some(num_token)); let response =
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
Ok(response) Ok(response)
} }
+3 -1
View File
@@ -63,6 +63,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
let mes_render = self.chat_template.apply_chat_template(&mes)?; let mes_render = self.chat_template.apply_chat_template(&mes)?;
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let mut generate = Vec::new(); let mut generate = Vec::new();
let sample_len = mes.max_tokens.unwrap_or(2048); let sample_len = mes.max_tokens.unwrap_or(2048);
@@ -81,7 +82,8 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
let num_token = generate.len() as u32; let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?; let res = self.tokenizer.token_decode(generate)?;
self.minicpm.clear_kv_cache(); self.minicpm.clear_kv_cache();
let response = build_completion_response(res, &self.model_name, Some(num_token)); let response =
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
Ok(response) Ok(response)
} }
fn generate_stream( fn generate_stream(
+3 -1
View File
@@ -71,6 +71,7 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
self.pre_processor.process_info(&mes, &mes_render)?; self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; let mut input_ids = self.tokenizer.text_encode(replace_text, &self.device)?;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let image_mask = get_equal_mask(&input_ids, self.cfg.image_token_id)?; let image_mask = get_equal_mask(&input_ids, self.cfg.image_token_id)?;
@@ -107,7 +108,8 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
let num_token = generate.len() as u32; let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?; let res = self.tokenizer.token_decode(generate)?;
self.paddleocr_vl.clear_kv_cache(); self.paddleocr_vl.clear_kv_cache();
let response = build_completion_response(res, &self.model_name, Some(num_token)); let response =
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
Ok(response) Ok(response)
} }
+3 -1
View File
@@ -75,6 +75,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
.tokenizer .tokenizer
.text_encode(input.replace_text.clone(), &self.device)?; .text_encode(input.replace_text.clone(), &self.device)?;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let mut pixel_values = input.pixel_values.as_ref(); let mut pixel_values = input.pixel_values.as_ref();
let image_grid_thw = input.image_grid_thw.as_ref(); let image_grid_thw = input.image_grid_thw.as_ref();
@@ -121,7 +122,8 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
let num_token = generate.len() as u32; let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?; let res = self.tokenizer.token_decode(generate)?;
self.qwen2_5_vl.clear_kv_cache(); self.qwen2_5_vl.clear_kv_cache();
let response = build_completion_response(res, &self.model_name, Some(num_token)); let response =
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
Ok(response) Ok(response)
} }
+15 -5
View File
@@ -11,8 +11,8 @@ use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig};
use crate::models::qwen3::model::Qwen3Model; use crate::models::qwen3::model::Qwen3Model;
// use crate::models::GenerateStream; // use crate::models::GenerateStream;
use crate::utils::{ use crate::utils::{
build_completion_chunk_response, build_completion_response, find_type_files, get_device, build_completion_chunk_response, build_completion_response, extract_metadata_value,
get_dtype, get_logit_processor, find_type_files, get_device, get_dtype, get_logit_processor,
}; };
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel}; use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
@@ -73,9 +73,14 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
}; };
let mut logit_processor = let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); 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::<bool>(&mes.metadata, "enable_thinking");
// let mes_render = self.chat_template.apply_chat_template(&mes)?;
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 input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let mut generate = Vec::new(); let mut generate = Vec::new();
let sample_len = mes.max_tokens.unwrap_or(2048); let sample_len = mes.max_tokens.unwrap_or(2048);
@@ -94,7 +99,8 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
let num_token = generate.len() as u32; let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?; let res = self.tokenizer.token_decode(generate)?;
self.qwen3.clear_kv_cache(); self.qwen3.clear_kv_cache();
let response = build_completion_response(res, &self.model_name, Some(num_token)); let response =
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
Ok(response) Ok(response)
} }
fn generate_stream( fn generate_stream(
@@ -123,7 +129,11 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
}; };
let mut logit_processor = let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); 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::<bool>(&mes.metadata, "enable_thinking");
// let mes_render = self.chat_template.apply_chat_template(&mes)?;
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 input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
+8 -2
View File
@@ -74,6 +74,7 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
.tokenizer .tokenizer
.text_encode(input.replace_text.clone(), &self.device)?; .text_encode(input.replace_text.clone(), &self.device)?;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let mut pixel_values = input.pixel_values.as_ref(); let mut pixel_values = input.pixel_values.as_ref();
let image_grid_thw = input.image_grid_thw.as_ref(); let image_grid_thw = input.image_grid_thw.as_ref();
@@ -102,10 +103,15 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
pixel_values = None; pixel_values = None;
pixel_values_video = None; pixel_values_video = None;
} }
let num_token = generate.len() as u32; let completion_tokens = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?; let res = self.tokenizer.token_decode(generate)?;
self.qwen3_5.clear_cache(); self.qwen3_5.clear_cache();
let response = build_completion_response(res, &self.model_name, Some(num_token)); let response = build_completion_response(
res,
&self.model_name,
Some(completion_tokens),
Some(prompt_tokens),
);
Ok(response) Ok(response)
} }
+4 -1
View File
@@ -90,10 +90,12 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
.process_info(&mes, &render_text, &self.tokenizer)?; .process_info(&mes, &render_text, &self.tokenizer)?;
let sample_len = mes.max_tokens.unwrap_or(1024); let sample_len = mes.max_tokens.unwrap_or(1024);
let mut generate = Vec::new(); let mut generate = Vec::new();
let mut prompt_tokens = 0u32;
for data in audio_datas.iter() { for data in audio_datas.iter() {
let mut input_ids = data.input_ids.clone(); let mut input_ids = data.input_ids.clone();
let mut input_features = Some(data.input_features.clone().to_dtype(self.dtype)?); let mut input_features = Some(data.input_features.clone().to_dtype(self.dtype)?);
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
prompt_tokens += seq_len as u32;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
for _ in 0..sample_len { for _ in 0..sample_len {
let logits = let logits =
@@ -114,7 +116,8 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
} }
let num_token = generate.len() as u32; let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?; let res = self.tokenizer.token_decode(generate)?;
let response = build_completion_response(res, &self.model_name, Some(num_token)); let response =
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
Ok(response) Ok(response)
} }
+15 -5
View File
@@ -16,8 +16,8 @@ use crate::{
}, },
tokenizer::TokenizerModel, tokenizer::TokenizerModel,
utils::{ utils::{
build_completion_chunk_response, build_completion_response, find_type_files, get_device, build_completion_chunk_response, build_completion_response, extract_metadata_value,
get_dtype, get_logit_processor, find_type_files, get_device, get_dtype, get_logit_processor,
}, },
}; };
@@ -80,12 +80,17 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
}; };
let mut logit_processor = let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); 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::<bool>(&mes.metadata, "enable_thinking");
// let mes_render = self.chat_template.apply_chat_template(&mes)?;
let mes_render = self
.chat_template
.apply_chat_temp_think(&mes, enable_thinking)?;
let input = self.pre_processor.process_info(&mes, &mes_render)?; let input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self let mut input_ids = self
.tokenizer .tokenizer
.text_encode(input.replace_text.clone(), &self.device)?; .text_encode(input.replace_text.clone(), &self.device)?;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let mut pixel_values = input.pixel_values.as_ref(); let mut pixel_values = input.pixel_values.as_ref();
let image_grid_thw = input.image_grid_thw.as_ref(); let image_grid_thw = input.image_grid_thw.as_ref();
@@ -120,7 +125,8 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
let num_token = generate.len() as u32; let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?; let res = self.tokenizer.token_decode(generate)?;
self.qwen3_vl.clear_kv_cache(); self.qwen3_vl.clear_kv_cache();
let response = build_completion_response(res, &self.model_name, Some(num_token)); let response =
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
Ok(response) Ok(response)
} }
@@ -150,7 +156,11 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
}; };
let mut logit_processor = let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); 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::<bool>(&mes.metadata, "enable_thinking");
// let mes_render = self.chat_template.apply_chat_template(&mes)?;
let mes_render = self
.chat_template
.apply_chat_temp_think(&mes, enable_thinking)?;
let input = self.pre_processor.process_info(&mes, &mes_render)?; let input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self let mut input_ids = self
.tokenizer .tokenizer
+12 -6
View File
@@ -479,16 +479,22 @@ pub fn build_audio_completion_response(
pub fn build_completion_response( pub fn build_completion_response(
res: String, res: String,
model_name: &str, model_name: &str,
num_tokens: Option<u32>, completion_tokens: Option<u32>,
prompt_tokens: Option<u32>,
) -> ChatCompletionResponse { ) -> ChatCompletionResponse {
let id = uuid::Uuid::new_v4().to_string(); let id = uuid::Uuid::new_v4().to_string();
let usage = num_tokens.map(|num| Usage { let usage = if prompt_tokens.is_none() && completion_tokens.is_none() {
prompt_tokens: None, None
completion_tokens: None, } else {
total_tokens: num, Some(Usage {
prompt_tokens,
completion_tokens,
total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0),
prompt_tokens_details: None, prompt_tokens_details: None,
completion_tokens_details: None, completion_tokens_details: None,
}); })
};
let mut response = ChatCompletionResponse { let mut response = ChatCompletionResponse {
id: Some(id), id: Some(id),
choices: vec![], choices: vec![],