From 037207eccd0f80632345338e8ccc7ed5d1b7e249 Mon Sep 17 00:00:00 2001 From: jhqxxx <18280426169@163.com> Date: Fri, 6 Mar 2026 19:08:58 +0800 Subject: [PATCH] update generate param init --- src/models/deepseek_ocr/generate.rs | 68 ++++++++--------------------- src/models/fun_asr_nano/generate.rs | 34 +++++---------- src/models/glm_asr_nano/generate.rs | 10 +---- src/models/hunyuan_ocr/generate.rs | 34 +++++---------- src/models/minicpm4/generate.rs | 10 +---- src/models/paddleocr_vl/generate.rs | 10 +---- src/models/qwen2_5vl/generate.rs | 10 +---- src/models/qwen3/generate.rs | 34 +++++---------- src/models/qwen3_5/generate.rs | 10 +---- src/models/qwen3_asr/generate.rs | 24 ++++------ src/models/qwen3vl/generate.rs | 34 +++++---------- tests/test_deepseek_ocr.rs | 2 +- 12 files changed, 77 insertions(+), 203 deletions(-) diff --git a/src/models/deepseek_ocr/generate.rs b/src/models/deepseek_ocr/generate.rs index edf5b71..f5ca3c4 100644 --- a/src/models/deepseek_ocr/generate.rs +++ b/src/models/deepseek_ocr/generate.rs @@ -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 { - 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::().unwrap_or(640); - if self.size.contains(&size) { size } else { 640 } + let base_size = extract_metadata_value::(&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::().unwrap_or(640); - if self.size.contains(&size) { size } else { 640 } + let image_size = extract_metadata_value::(&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::().unwrap_or(false) - } else { - false - }; - let seed = match mes.seed { - None => 34562u64, - Some(s) => s as u64, - }; + let crop_mode = extract_metadata_value::(&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::().unwrap_or(640); - if self.size.contains(&size) { size } else { 640 } + let base_size = extract_metadata_value::(&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::().unwrap_or(640); - if self.size.contains(&size) { size } else { 640 } + let image_size = extract_metadata_value::(&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::().unwrap_or(false) - } else { - false - }; - let seed = match mes.seed { - None => 34562u64, - Some(s) => s as u64, - }; + let crop_mode = extract_metadata_value::(&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 diff --git a/src/models/fun_asr_nano/generate.rs b/src/models/fun_asr_nano/generate.rs index d10ea9d..a192826 100644 --- a/src/models/fun_asr_nano/generate.rs +++ b/src/models/fun_asr_nano/generate.rs @@ -94,19 +94,12 @@ impl FunAsrNanoGenerateModel { impl GenerateModel for FunAsrNanoGenerateModel { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - 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)?; diff --git a/src/models/glm_asr_nano/generate.rs b/src/models/glm_asr_nano/generate.rs index 35e7918..ca723fb 100644 --- a/src/models/glm_asr_nano/generate.rs +++ b/src/models/glm_asr_nano/generate.rs @@ -65,10 +65,7 @@ impl<'a> GlmAsrNanoGenerateModel<'a> { impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - 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) = diff --git a/src/models/hunyuan_ocr/generate.rs b/src/models/hunyuan_ocr/generate.rs index e600626..0bc7539 100644 --- a/src/models/hunyuan_ocr/generate.rs +++ b/src/models/hunyuan_ocr/generate.rs @@ -68,19 +68,12 @@ impl<'a> HunyuanOCRGenerateModel<'a> { impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - 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)?; diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs index faf1027..46cf783 100644 --- a/src/models/minicpm4/generate.rs +++ b/src/models/minicpm4/generate.rs @@ -55,10 +55,7 @@ impl<'a> MiniCPMGenerateModel<'a> { impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - 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)?; diff --git a/src/models/paddleocr_vl/generate.rs b/src/models/paddleocr_vl/generate.rs index b0a7e99..f809358 100644 --- a/src/models/paddleocr_vl/generate.rs +++ b/src/models/paddleocr_vl/generate.rs @@ -61,10 +61,7 @@ impl<'a> PaddleOCRVLGenerateModel<'a> { impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - 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) = diff --git a/src/models/qwen2_5vl/generate.rs b/src/models/qwen2_5vl/generate.rs index 1e93afa..486c0ce 100644 --- a/src/models/qwen2_5vl/generate.rs +++ b/src/models/qwen2_5vl/generate.rs @@ -64,10 +64,7 @@ impl<'a> Qwen2_5VLGenerateModel<'a> { impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - 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)?; diff --git a/src/models/qwen3/generate.rs b/src/models/qwen3/generate.rs index 5ca3775..3f5d8ec 100644 --- a/src/models/qwen3/generate.rs +++ b/src/models/qwen3/generate.rs @@ -58,19 +58,12 @@ impl<'a> Qwen3GenerateModel<'a> { impl<'a> GenerateModel for Qwen3GenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - 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::(&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::(&mes.metadata, "enable_thinking"); diff --git a/src/models/qwen3_5/generate.rs b/src/models/qwen3_5/generate.rs index d0d990e..82b23ed 100644 --- a/src/models/qwen3_5/generate.rs +++ b/src/models/qwen3_5/generate.rs @@ -60,10 +60,7 @@ impl<'a> Qwen3_5GenerateModel<'a> { impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - 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::(&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::(&mes.metadata, "enable_thinking"); let mes_render = self diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs index fef1cd3..d366411 100644 --- a/src/models/qwen3_asr/generate.rs +++ b/src/models/qwen3_asr/generate.rs @@ -75,14 +75,10 @@ impl<'a> Qwen3AsrGenerateModel<'a> { impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - 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 diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs index de3c0d3..90da57d 100644 --- a/src/models/qwen3vl/generate.rs +++ b/src/models/qwen3vl/generate.rs @@ -65,19 +65,12 @@ impl<'a> Qwen3VLGenerateModel<'a> { impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - 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::(&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::(&mes.metadata, "enable_thinking"); diff --git a/tests/test_deepseek_ocr.rs b/tests/test_deepseek_ocr.rs index d716ce0..10e9bf5 100644 --- a/tests/test_deepseek_ocr.rs +++ b/tests/test_deepseek_ocr.rs @@ -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#" {