update generate param init

This commit is contained in:
jhqxxx
2026-03-06 19:08:58 +08:00
parent 548eb8185e
commit 037207eccd
12 changed files with 77 additions and 203 deletions
+18 -50
View File
@@ -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,
},
};
@@ -62,36 +62,20 @@ impl DeepseekOCRGenerateModel {
impl GenerateModel for DeepseekOCRGenerateModel {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let base_size = if let Some(map) = &mes.metadata
&& map.contains_key("base_size")
{
let size = map.get("base_size").unwrap();
let size = size.parse::<u32>().unwrap_or(640);
if self.size.contains(&size) { size } else { 640 }
let base_size = extract_metadata_value::<u32>(&mes.metadata, "base_size").unwrap_or(640);
let base_size = if self.size.contains(&base_size) {
base_size
} else {
640
};
let image_size = if let Some(map) = &mes.metadata
&& map.contains_key("image_size")
{
let size = map.get("image_size").unwrap();
let size = size.parse::<u32>().unwrap_or(640);
if self.size.contains(&size) { size } else { 640 }
let image_size = extract_metadata_value::<u32>(&mes.metadata, "image_size").unwrap_or(640);
let image_size = if self.size.contains(&image_size) {
image_size
} else {
640
};
let crop_mode = if let Some(map) = &mes.metadata
&& map.contains_key("crop_mode")
{
let size = map.get("crop_mode").unwrap();
size.parse::<bool>().unwrap_or(false)
} else {
false
};
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
let crop_mode = extract_metadata_value::<bool>(&mes.metadata, "crop_mode").unwrap_or(false);
let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
let (mut input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self
.processor
@@ -147,36 +131,20 @@ impl GenerateModel for DeepseekOCRGenerateModel {
+ '_,
>,
> {
let base_size = if let Some(map) = &mes.metadata
&& map.contains_key("base_size")
{
let size = map.get("base_size").unwrap();
let size = size.parse::<u32>().unwrap_or(640);
if self.size.contains(&size) { size } else { 640 }
let base_size = extract_metadata_value::<u32>(&mes.metadata, "base_size").unwrap_or(640);
let base_size = if self.size.contains(&base_size) {
base_size
} else {
640
};
let image_size = if let Some(map) = &mes.metadata
&& map.contains_key("image_size")
{
let size = map.get("image_size").unwrap();
let size = size.parse::<u32>().unwrap_or(640);
if self.size.contains(&size) { size } else { 640 }
let image_size = extract_metadata_value::<u32>(&mes.metadata, "image_size").unwrap_or(640);
let image_size = if self.size.contains(&image_size) {
image_size
} else {
640
};
let crop_mode = if let Some(map) = &mes.metadata
&& map.contains_key("crop_mode")
{
let size = map.get("crop_mode").unwrap();
size.parse::<bool>().unwrap_or(false)
} else {
false
};
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
let crop_mode = extract_metadata_value::<bool>(&mes.metadata, "crop_mode").unwrap_or(false);
let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
let (mut input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self
.processor
+10 -24
View File
@@ -94,19 +94,12 @@ impl FunAsrNanoGenerateModel {
impl GenerateModel for FunAsrNanoGenerateModel {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let top_p = match mes.top_p {
None => self.generation_config.top_p,
Some(top_p) => top_p,
};
let temperature = mes
.temperature
.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 = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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 (speech, fbank_mask, mut input_ids) =
@@ -156,19 +149,12 @@ impl GenerateModel for FunAsrNanoGenerateModel {
+ '_,
>,
> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let top_p = match mes.top_p {
None => self.generation_config.top_p,
Some(top_p) => top_p,
};
let temperature = mes
.temperature
.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 = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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 (speech, fbank_mask, input_ids) = self.processor.process_info(&mes, &self.tokenizer)?;
+2 -8
View File
@@ -65,10 +65,7 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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) =
@@ -122,10 +119,7 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
+ '_,
>,
> {
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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) =
+10 -24
View File
@@ -68,19 +68,12 @@ impl<'a> HunyuanOCRGenerateModel<'a> {
impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let top_p = match mes.top_p {
None => self.generation_config.top_p,
Some(top_p) => top_p,
};
let temperature = mes
.temperature
.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 = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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)?;
@@ -139,19 +132,12 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
+ '_,
>,
> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let top_p = match mes.top_p {
None => self.generation_config.top_p,
Some(top_p) => top_p,
};
let temperature = mes
.temperature
.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 = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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)?;
+2 -8
View File
@@ -55,10 +55,7 @@ impl<'a> MiniCPMGenerateModel<'a> {
impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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)?;
@@ -97,10 +94,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
+ '_,
>,
> {
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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)?;
+2 -8
View File
@@ -61,10 +61,7 @@ impl<'a> PaddleOCRVLGenerateModel<'a> {
impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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) =
@@ -124,10 +121,7 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
+ '_,
>,
> {
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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) =
+2 -8
View File
@@ -64,10 +64,7 @@ impl<'a> Qwen2_5VLGenerateModel<'a> {
impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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)?;
@@ -138,10 +135,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
+ '_,
>,
> {
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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)?;
+10 -24
View File
@@ -58,19 +58,12 @@ impl<'a> Qwen3GenerateModel<'a> {
impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let top_p = match mes.top_p {
None => self.generation_config.top_p,
Some(top_p) => top_p,
};
let temperature = mes
.temperature
.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 = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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 enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
@@ -114,19 +107,12 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
+ '_,
>,
> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let top_p = match mes.top_p {
None => self.generation_config.top_p,
Some(top_p) => top_p,
};
let temperature = mes
.temperature
.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 = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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 enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
+2 -8
View File
@@ -60,10 +60,7 @@ impl<'a> Qwen3_5GenerateModel<'a> {
impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
let seed = mes.seed.unwrap_or(34562) as u64;
let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
let mes_render = self
@@ -126,10 +123,7 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
+ '_,
>,
> {
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
let mes_render = self
+8 -16
View File
@@ -75,14 +75,10 @@ impl<'a> Qwen3AsrGenerateModel<'a> {
impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
let temperature = mes
.temperature
.unwrap_or(self.generation_config.temperature);
let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = get_logit_processor(Some(temperature), mes.top_p, None, seed);
let render_text = self.chat_template.apply_chat_template(&mes)?;
let audio_datas = self
@@ -132,14 +128,10 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
+ '_,
>,
> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
let temperature = mes
.temperature
.unwrap_or(self.generation_config.temperature);
let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = get_logit_processor(Some(temperature), mes.top_p, None, seed);
let render_text = self.chat_template.apply_chat_template(&mes)?;
let audio_datas = self
+10 -24
View File
@@ -65,19 +65,12 @@ impl<'a> Qwen3VLGenerateModel<'a> {
impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let top_p = match mes.top_p {
None => self.generation_config.top_p,
Some(top_p) => top_p,
};
let temperature = mes
.temperature
.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 = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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 enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
@@ -141,19 +134,12 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
+ '_,
>,
> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let top_p = match mes.top_p {
None => self.generation_config.top_p,
Some(top_p) => top_p,
};
let temperature = mes
.temperature
.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 = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
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 enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
+1 -1
View File
@@ -56,7 +56,7 @@ fn deepseek_ocr_generate() -> Result<()> {
#[tokio::test]
async fn deepseek_ocr_stream() -> Result<()> {
// test with cuda: RUST_BACKTRACE=1 cargo test -F cuda deepseek_ocr_stream -r -- --nocapture
// test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_deepseek_ocr deepseek_ocr_stream -r -- --nocapture
let message = r#"
{