From 42c95da0e392578ced192db9e3336152df13f891 Mon Sep 17 00:00:00 2001 From: YageGeng Date: Thu, 11 Dec 2025 10:50:35 +0800 Subject: [PATCH] refactor: remove unstable Option<&T> branch by returning owned Tensor --- src/models/deepseek_ocr/generate.rs | 2 +- src/models/deepseek_ocr/model.rs | 7 ++++--- src/models/hunyuan_ocr/generate.rs | 2 +- src/models/hunyuan_ocr/model.rs | 17 ++++------------- src/models/minicpm4/generate.rs | 2 +- src/models/minicpm4/model.rs | 19 ++++++++++++------- src/models/paddleocr_vl/generate.rs | 2 +- src/models/paddleocr_vl/model.rs | 8 ++++---- src/models/qwen2_5vl/generate.rs | 2 +- src/models/qwen2_5vl/model.rs | 8 ++++---- src/models/qwen2_5vl/processor.rs | 4 ++-- src/models/qwen3vl/generate.rs | 2 +- src/models/qwen3vl/model.rs | 16 ++++++++++------ src/models/qwen3vl/processor.rs | 10 ++++------ src/models/voxcpm/minicpm4.rs | 19 ++++++++++++------- src/models/voxcpm/tokenizer.rs | 4 ++-- 16 files changed, 64 insertions(+), 60 deletions(-) diff --git a/src/models/deepseek_ocr/generate.rs b/src/models/deepseek_ocr/generate.rs index 3904758..47f52e3 100644 --- a/src/models/deepseek_ocr/generate.rs +++ b/src/models/deepseek_ocr/generate.rs @@ -205,7 +205,7 @@ impl GenerateModel for DeepseekOCRGenerateModel { decode_ids.extend_from_slice(&error_tokens); } decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{}", e)))?; + let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; if decoded_token.contains("�") { error_tokens.push(next_token); if error_tokens.len() > 3 { diff --git a/src/models/deepseek_ocr/model.rs b/src/models/deepseek_ocr/model.rs index 2157498..6919afb 100644 --- a/src/models/deepseek_ocr/model.rs +++ b/src/models/deepseek_ocr/model.rs @@ -1072,16 +1072,17 @@ impl DeepseekV2Model { pub fn forward(&mut self, xs: &Tensor, seqlen_offset: usize) -> Result { let (bs, seq_len, _) = xs.dims3()?; let (cos, sin) = self.rope.forward(seqlen_offset, seq_len, xs.device())?; - let attention_mask: Option<&Tensor> = { + + let attention_mask: Option = { if seq_len <= 1 { None } else { - Some(&prepare_causal_attention_mask(bs, seq_len, 0, xs.device())?) + Some(prepare_causal_attention_mask(bs, seq_len, 0, xs.device())?) } }; let mut xs = xs.clone(); for layer in &mut self.layers { - xs = layer.forward(&xs, &cos, &sin, attention_mask)?; + xs = layer.forward(&xs, &cos, &sin, attention_mask.as_ref())?; } let xs = self.norm.forward(&xs)?; Ok(xs) diff --git a/src/models/hunyuan_ocr/generate.rs b/src/models/hunyuan_ocr/generate.rs index 04295d9..4a6de5f 100644 --- a/src/models/hunyuan_ocr/generate.rs +++ b/src/models/hunyuan_ocr/generate.rs @@ -183,7 +183,7 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> { decode_ids.extend_from_slice(&error_tokens); } decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{}", e)))?; + let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; if decoded_token.contains("�") { error_tokens.push(next_token); if error_tokens.len() > 3 { diff --git a/src/models/hunyuan_ocr/model.rs b/src/models/hunyuan_ocr/model.rs index a35c615..67492b9 100644 --- a/src/models/hunyuan_ocr/model.rs +++ b/src/models/hunyuan_ocr/model.rs @@ -482,20 +482,11 @@ impl HunYuanVLTextModel { ) -> Result { let (b_size, seq_len, _) = inputs_embeds.dims3()?; - // let position_ids = match position_ids { - // Some(ids) => ids.clone(), - // None => Tensor::arange( - // seqlen_offset as u32, - // (seq_len + seqlen_offset) as u32, - // inputs_embeds.device(), - // )? - // .unsqueeze(0)?, - // }; - let attention_mask: Option<&Tensor> = { + let attention_mask: Option = { if seq_len <= 1 { None } else { - Some(&prepare_causal_attention_mask( + Some(prepare_causal_attention_mask( b_size, seq_len, 0, @@ -514,9 +505,9 @@ impl HunYuanVLTextModel { { let (cos, sin) = get_xd_cos_sin(&cos, &sin, position_ids, self.xdrope_section.clone())?; - xs = layer.forward(&xs, &cos, &sin, attention_mask)?; + xs = layer.forward(&xs, &cos, &sin, attention_mask.as_ref())?; } else { - xs = layer.forward(&xs, &cos, &sin, attention_mask)?; + xs = layer.forward(&xs, &cos, &sin, attention_mask.as_ref())?; } } let xs = self.norm.forward(&xs)?; diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs index 100cdcf..45839f8 100644 --- a/src/models/minicpm4/generate.rs +++ b/src/models/minicpm4/generate.rs @@ -119,7 +119,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { decode_ids.extend_from_slice(&error_tokens); } decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{}", e)))?; + let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; if decoded_token.contains("�") { error_tokens.push(next_token); if error_tokens.len() > 3 { diff --git a/src/models/minicpm4/model.rs b/src/models/minicpm4/model.rs index cf48c5f..4a9ac61 100644 --- a/src/models/minicpm4/model.rs +++ b/src/models/minicpm4/model.rs @@ -232,11 +232,11 @@ impl MiniCPMModel { .embed_tokens .forward(input_ids)? .affine(self.cfg.scale_emb, 0.0)?; - let attention_mask: Option<&Tensor> = { + let attention_mask: Option = { if seq_len <= 1 { None } else { - Some(&prepare_causal_attention_mask( + Some(prepare_causal_attention_mask( bs, seq_len, 0, @@ -248,7 +248,8 @@ impl MiniCPMModel { let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; let mut hidden_states = input_embeds; for decode_layer in &self.layers { - hidden_states = decode_layer.forward(&hidden_states, &cos, &sin, attention_mask)?; + hidden_states = + decode_layer.forward(&hidden_states, &cos, &sin, attention_mask.as_ref())?; } hidden_states = self.norm.forward(&hidden_states)?; let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?; @@ -266,11 +267,11 @@ impl MiniCPMModel { .embed_tokens .forward(input_ids)? .affine(self.cfg.scale_emb, 0.0)?; - let attention_mask: Option<&Tensor> = { + let attention_mask: Option = { if seq_len <= 1 { None } else { - Some(&prepare_causal_attention_mask( + Some(prepare_causal_attention_mask( bs, seq_len, 0, @@ -281,8 +282,12 @@ impl MiniCPMModel { let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; let mut hidden_states = input_embeds; for decode_layer in &mut self.layers { - hidden_states = - decode_layer.forward_with_cache(&hidden_states, &cos, &sin, attention_mask)?; + hidden_states = decode_layer.forward_with_cache( + &hidden_states, + &cos, + &sin, + attention_mask.as_ref(), + )?; } hidden_states = self.norm.forward(&hidden_states)?; let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?; diff --git a/src/models/paddleocr_vl/generate.rs b/src/models/paddleocr_vl/generate.rs index 36e470d..a8d4234 100644 --- a/src/models/paddleocr_vl/generate.rs +++ b/src/models/paddleocr_vl/generate.rs @@ -162,7 +162,7 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> { decode_ids.extend_from_slice(&error_tokens); } decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{}", e)))?; + let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; if decoded_token.contains("�") { error_tokens.push(next_token); if error_tokens.len() > 3 { diff --git a/src/models/paddleocr_vl/model.rs b/src/models/paddleocr_vl/model.rs index ebfb86b..4f68f68 100644 --- a/src/models/paddleocr_vl/model.rs +++ b/src/models/paddleocr_vl/model.rs @@ -376,11 +376,11 @@ impl Ernie4_5Model { self.rope_scaling.mrope_section.clone(), )?; let mut xs = inputs_embeds.clone(); - let attention_mask: Option<&Tensor> = { + let attention_mask: Option = { if seq_len <= 1 { None } else { - Some(&prepare_causal_attention_mask( + Some(prepare_causal_attention_mask( b_size, seq_len, 0, @@ -389,7 +389,7 @@ impl Ernie4_5Model { } }; for layer in self.layers.iter_mut() { - xs = layer.forward(&xs, &cos, &sin, attention_mask)?; + xs = layer.forward(&xs, &cos, &sin, attention_mask.as_ref())?; } let xs = xs.apply(&self.norm)?; Ok(xs) @@ -564,7 +564,7 @@ impl PaddleOCRVLModel { } } Err(e) => { - println!("get vision_indices err: {}", e); + println!("get vision_indices err: {e}"); } }; diff --git a/src/models/qwen2_5vl/generate.rs b/src/models/qwen2_5vl/generate.rs index 63c9376..c7e7d14 100644 --- a/src/models/qwen2_5vl/generate.rs +++ b/src/models/qwen2_5vl/generate.rs @@ -189,7 +189,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { decode_ids.extend_from_slice(&error_tokens); } decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{}", e)))?; + let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; if decoded_token.contains("�") { error_tokens.push(next_token); if error_tokens.len() > 3 { diff --git a/src/models/qwen2_5vl/model.rs b/src/models/qwen2_5vl/model.rs index 457f100..0260532 100644 --- a/src/models/qwen2_5vl/model.rs +++ b/src/models/qwen2_5vl/model.rs @@ -713,11 +713,11 @@ impl Qwen2_5VLTextModel { self.rope_scaling.mrope_section.clone(), )?; let mut xs = inputs_embeds.clone(); - let attention_mask: Option<&Tensor> = { + let attention_mask: Option = { if seq_len <= 1 { None } else { - Some(&prepare_causal_attention_mask( + Some(prepare_causal_attention_mask( b_size, seq_len, 0, @@ -726,7 +726,7 @@ impl Qwen2_5VLTextModel { } }; for layer in self.layers.iter_mut() { - xs = layer.forward(&xs, &cos, &sin, attention_mask)?; + xs = layer.forward(&xs, &cos, &sin, attention_mask.as_ref())?; } let xs = xs.apply(&self.norm)?; Ok(xs) @@ -898,7 +898,7 @@ impl Qwen2_5VLModel { } } Err(e) => { - println!("get vision_indices err: {}", e); + println!("get vision_indices err: {e}"); } }; diff --git a/src/models/qwen2_5vl/processor.rs b/src/models/qwen2_5vl/processor.rs index 53e3da8..babd8a3 100644 --- a/src/models/qwen2_5vl/processor.rs +++ b/src/models/qwen2_5vl/processor.rs @@ -243,7 +243,7 @@ impl Qwen2_5VLProcessor { let image = get_image(file); match image { Ok(img) => file_vec.push(img), - Err(e) => println!("get_image err: {:?}", e), + Err(e) => println!("get_image err: {e:?}"), }; } if !file_vec.is_empty() { @@ -253,7 +253,7 @@ impl Qwen2_5VLProcessor { pixel_values = Some(img_input.data); image_grid_thw = Some(img_input.grid_thw); } - Err(e) => println!("img process_images err: {:?}", e), + Err(e) => println!("img process_images err: {e:?}"), }; } } diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs index 25185b8..fd7cedc 100644 --- a/src/models/qwen3vl/generate.rs +++ b/src/models/qwen3vl/generate.rs @@ -191,7 +191,7 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { decode_ids.extend_from_slice(&error_tokens); } decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{}", e)))?; + let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; if decoded_token.contains("�") { error_tokens.push(next_token); if error_tokens.len() > 3 { diff --git a/src/models/qwen3vl/model.rs b/src/models/qwen3vl/model.rs index a84610b..8872893 100644 --- a/src/models/qwen3vl/model.rs +++ b/src/models/qwen3vl/model.rs @@ -747,11 +747,11 @@ impl Qwen3VLTextModel { self.mrope_section.clone(), )?; let mut xs = inputs_embeds.clone(); - let attention_mask: Option<&Tensor> = { + let attention_mask: Option = { if seq_len <= 1 { None } else { - Some(&prepare_causal_attention_mask( + Some(prepare_causal_attention_mask( b_size, seq_len, 0, @@ -760,7 +760,7 @@ impl Qwen3VLTextModel { } }; for (layer_idx, layer) in self.layers.iter_mut().enumerate() { - xs = layer.forward(&xs, &cos, &sin, attention_mask)?; + xs = layer.forward(&xs, &cos, &sin, attention_mask.as_ref())?; if let Some(deepstack_embeds) = deepstack_visual_embeds.as_ref() && layer_idx < deepstack_embeds.len() { @@ -867,7 +867,7 @@ impl Qwen3VLModel { .reshape((*t as usize, ()))?, ); } - Some(&Tensor::cat(&v_thw_vec, 0)?) + Some(Tensor::cat(&v_thw_vec, 0)?) } None => None, }; @@ -918,7 +918,11 @@ impl Qwen3VLModel { text_end = vision_indices_vec[j]; } if token == video_token_id as u32 { - thw = video_grid_thw.unwrap().i(video_index)?.to_vec1::()?; + thw = video_grid_thw + .as_ref() + .unwrap() + .i(video_index)? + .to_vec1::()?; text_end = vision_indices_vec[j]; video_index += 1; } @@ -987,7 +991,7 @@ impl Qwen3VLModel { } } Err(e) => { - println!("get vision_indices err: {}", e); + println!("get vision_indices err: {e}"); } }; if text_start < input_ids_i.dim(0)? as u32 { diff --git a/src/models/qwen3vl/processor.rs b/src/models/qwen3vl/processor.rs index 1b73058..f69f75c 100644 --- a/src/models/qwen3vl/processor.rs +++ b/src/models/qwen3vl/processor.rs @@ -310,7 +310,7 @@ impl Qwen3VLProcessor { let image = get_image(file); match image { Ok(img) => file_vec.push(img), - Err(e) => println!("get_image err: {:?}", e), + Err(e) => println!("get_image err: {e:?}"), }; } if !file_vec.is_empty() { @@ -320,7 +320,7 @@ impl Qwen3VLProcessor { pixel_values = Some(img_input.data); image_grid_thw = Some(img_input.grid_thw); } - Err(e) => println!("img process_images err: {:?}", e), + Err(e) => println!("img process_images err: {e:?}"), }; } } @@ -435,14 +435,12 @@ pub fn video_smart_resize( ) -> Result<(u32, u32)> { if num_frames < temporal_factor { return Err(anyhow!(format!( - "{} must be larger than temporal_factor {}", - num_frames, temporal_factor + "{num_frames} must be larger than temporal_factor {temporal_factor}" ))); } if height < factor || width < factor { return Err(anyhow!(format!( - "height:{} or width:{} must be larger than factor:{}", - height, width, factor + "height:{height} or width:{width} must be larger than factor:{factor}" ))); } if std::cmp::max(height, width) / std::cmp::min(height, width) > 200 { diff --git a/src/models/voxcpm/minicpm4.rs b/src/models/voxcpm/minicpm4.rs index 8ee90a7..45529a4 100644 --- a/src/models/voxcpm/minicpm4.rs +++ b/src/models/voxcpm/minicpm4.rs @@ -266,11 +266,11 @@ impl MiniCPMModel { is_causal: bool, ) -> Result { let (bs, seq_len, _) = input_embeds.dims3()?; - let attention_mask: Option<&Tensor> = { + let attention_mask: Option = { if !is_causal || seq_len <= 1 { None } else { - Some(&prepare_causal_attention_mask( + Some(prepare_causal_attention_mask( bs, seq_len, 0, @@ -281,7 +281,8 @@ impl MiniCPMModel { let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; let mut hidden_states = input_embeds.clone(); for decode_layer in &self.layers { - hidden_states = decode_layer.forward(&hidden_states, &cos, &sin, attention_mask)?; + hidden_states = + decode_layer.forward(&hidden_states, &cos, &sin, attention_mask.as_ref())?; } hidden_states = self.norm.forward(&hidden_states)?; Ok(hidden_states) @@ -298,11 +299,11 @@ impl MiniCPMModel { _ => return Err(anyhow!("MiniCPMModelinput_embeds illigal")), }; let (bs, seq_len, _) = input_embeds.dims3()?; - let attention_mask: Option<&Tensor> = { + let attention_mask: Option = { if seq_len <= 1 { None } else { - Some(&prepare_causal_attention_mask( + Some(prepare_causal_attention_mask( bs, seq_len, 0, @@ -313,8 +314,12 @@ impl MiniCPMModel { let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; let mut hidden_states = input_embeds.clone(); for decode_layer in &mut self.layers { - hidden_states = - decode_layer.forward_with_cache(&hidden_states, &cos, &sin, attention_mask)?; + hidden_states = decode_layer.forward_with_cache( + &hidden_states, + &cos, + &sin, + attention_mask.as_ref(), + )?; } hidden_states = self.norm.forward(&hidden_states)?; diff --git a/src/models/voxcpm/tokenizer.rs b/src/models/voxcpm/tokenizer.rs index 074662f..1e1e5a0 100644 --- a/src/models/voxcpm/tokenizer.rs +++ b/src/models/voxcpm/tokenizer.rs @@ -19,7 +19,7 @@ impl SingleChineseTokenizer { "tokenizer.json not exists in model path" ); let tokenizer = Tokenizer::from_file(tokenizer_file) - .map_err(|e| anyhow!(format!("tokenizer from file error{}", e)))?; + .map_err(|e| anyhow!(format!("tokenizer from file error{e}")))?; let mut multichar_tokens = Vec::new(); for (token, _) in tokenizer.get_vocab(false) { let len = token.chars().count(); @@ -42,7 +42,7 @@ impl SingleChineseTokenizer { let encode = self .tokenizer .encode(text, false) - .map_err(|e| anyhow!(format!("tokenizer encode error: {}", e)))?; + .map_err(|e| anyhow!(format!("tokenizer encode error: {e}")))?; let tokens = encode.get_tokens(); // println!("tokens: {:?}", tokens); let mut split_character = Vec::new();