usage add prompt_tokens
This commit is contained in:
@@ -102,6 +102,7 @@ impl GenerateModel for DeepseekOCRGenerateModel {
|
||||
let mut images_spatial_crop_t = Some(&images_spatial_crop_t);
|
||||
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(1024);
|
||||
for _ in 0..sample_len {
|
||||
@@ -130,7 +131,8 @@ impl GenerateModel for DeepseekOCRGenerateModel {
|
||||
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, Some(num_token));
|
||||
let response =
|
||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
|
||||
@@ -114,6 +114,7 @@ impl GenerateModel for FunAsrNanoGenerateModel {
|
||||
let mut speech = Some(speech.to_dtype(self.dtype)?);
|
||||
let mut fbank_mask = Some(&fbank_mask);
|
||||
let mut seq_len = input_ids.dim(1)?;
|
||||
let prompt_tokens = seq_len as u32;
|
||||
let mut seqlen_offset = 0;
|
||||
let mut generate = Vec::new();
|
||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||
@@ -139,7 +140,8 @@ impl GenerateModel for FunAsrNanoGenerateModel {
|
||||
let num_token = generate.len() as u32;
|
||||
let res = self.tokenizer.token_decode(generate)?;
|
||||
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)
|
||||
}
|
||||
|
||||
|
||||
@@ -77,6 +77,7 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
|
||||
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<u32> = Vec::new();
|
||||
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 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));
|
||||
let response =
|
||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
|
||||
@@ -93,6 +93,7 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
|
||||
let mut pixel_values = data.pixel_values;
|
||||
let mut image_grid_thw = data.image_grid_thw;
|
||||
let mut seq_len = input_ids.dim(1)?;
|
||||
let prompt_tokens = seq_len as u32;
|
||||
let mut seqlen_offset = 0;
|
||||
let mut generate: Vec<u32> = Vec::new();
|
||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||
@@ -122,7 +123,8 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
|
||||
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));
|
||||
let response =
|
||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
|
||||
@@ -63,6 +63,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
||||
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 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 res = self.tokenizer.token_decode(generate)?;
|
||||
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)
|
||||
}
|
||||
fn generate_stream(
|
||||
|
||||
@@ -71,6 +71,7 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
|
||||
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 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 res = self.tokenizer.token_decode(generate)?;
|
||||
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)
|
||||
}
|
||||
|
||||
|
||||
@@ -75,6 +75,7 @@ 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 seqlen_offset = 0;
|
||||
let mut pixel_values = input.pixel_values.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 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));
|
||||
let response =
|
||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
|
||||
@@ -11,8 +11,8 @@ 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, get_logit_processor,
|
||||
build_completion_chunk_response, build_completion_response, extract_metadata_value,
|
||||
find_type_files, get_device, get_dtype, get_logit_processor,
|
||||
};
|
||||
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
||||
|
||||
@@ -73,9 +73,14 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
|
||||
};
|
||||
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::<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 seq_len = input_ids.dim(1)?;
|
||||
let prompt_tokens = seq_len as u32;
|
||||
let mut seqlen_offset = 0;
|
||||
let mut generate = Vec::new();
|
||||
let sample_len = mes.max_tokens.unwrap_or(2048);
|
||||
@@ -94,7 +99,8 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
|
||||
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));
|
||||
let response =
|
||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
||||
Ok(response)
|
||||
}
|
||||
fn generate_stream(
|
||||
@@ -123,7 +129,11 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
|
||||
};
|
||||
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::<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 seq_len = input_ids.dim(1)?;
|
||||
let mut seqlen_offset = 0;
|
||||
|
||||
@@ -74,6 +74,7 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'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 seqlen_offset = 0;
|
||||
let mut pixel_values = input.pixel_values.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_video = None;
|
||||
}
|
||||
let num_token = generate.len() as u32;
|
||||
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(num_token));
|
||||
let response = build_completion_response(
|
||||
res,
|
||||
&self.model_name,
|
||||
Some(completion_tokens),
|
||||
Some(prompt_tokens),
|
||||
);
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
|
||||
@@ -90,10 +90,12 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
|
||||
.process_info(&mes, &render_text, &self.tokenizer)?;
|
||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||
let mut generate = Vec::new();
|
||||
let mut prompt_tokens = 0u32;
|
||||
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 logits =
|
||||
@@ -114,7 +116,8 @@ 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));
|
||||
let response =
|
||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
|
||||
@@ -16,8 +16,8 @@ use crate::{
|
||||
},
|
||||
tokenizer::TokenizerModel,
|
||||
utils::{
|
||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||
get_dtype, get_logit_processor,
|
||||
build_completion_chunk_response, build_completion_response, extract_metadata_value,
|
||||
find_type_files, get_device, get_dtype, get_logit_processor,
|
||||
},
|
||||
};
|
||||
|
||||
@@ -80,12 +80,17 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
||||
};
|
||||
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::<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 mut 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();
|
||||
@@ -120,7 +125,8 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
||||
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));
|
||||
let response =
|
||||
build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens));
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
@@ -150,7 +156,11 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
||||
};
|
||||
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::<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 mut input_ids = self
|
||||
.tokenizer
|
||||
|
||||
Reference in New Issue
Block a user