refactor: remove unstable Option<&T> branch by returning owned Tensor

This commit is contained in:
YageGeng
2025-12-11 10:50:35 +08:00
parent 0dc2840b0e
commit 42c95da0e3
16 changed files with 64 additions and 60 deletions
+1 -1
View File
@@ -205,7 +205,7 @@ impl GenerateModel for DeepseekOCRGenerateModel {
decode_ids.extend_from_slice(&error_tokens); decode_ids.extend_from_slice(&error_tokens);
} }
decode_ids.push(next_token); 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("") { if decoded_token.contains("") {
error_tokens.push(next_token); error_tokens.push(next_token);
if error_tokens.len() > 3 { if error_tokens.len() > 3 {
+4 -3
View File
@@ -1072,16 +1072,17 @@ impl DeepseekV2Model {
pub fn forward(&mut self, xs: &Tensor, seqlen_offset: usize) -> Result<Tensor> { pub fn forward(&mut self, xs: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
let (bs, seq_len, _) = xs.dims3()?; let (bs, seq_len, _) = xs.dims3()?;
let (cos, sin) = self.rope.forward(seqlen_offset, seq_len, xs.device())?; let (cos, sin) = self.rope.forward(seqlen_offset, seq_len, xs.device())?;
let attention_mask: Option<&Tensor> = {
let attention_mask: Option<Tensor> = {
if seq_len <= 1 { if seq_len <= 1 {
None None
} else { } 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(); let mut xs = xs.clone();
for layer in &mut self.layers { 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)?; let xs = self.norm.forward(&xs)?;
Ok(xs) Ok(xs)
+1 -1
View File
@@ -183,7 +183,7 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
decode_ids.extend_from_slice(&error_tokens); decode_ids.extend_from_slice(&error_tokens);
} }
decode_ids.push(next_token); 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("") { if decoded_token.contains("") {
error_tokens.push(next_token); error_tokens.push(next_token);
if error_tokens.len() > 3 { if error_tokens.len() > 3 {
+4 -13
View File
@@ -482,20 +482,11 @@ impl HunYuanVLTextModel {
) -> Result<Tensor> { ) -> Result<Tensor> {
let (b_size, seq_len, _) = inputs_embeds.dims3()?; let (b_size, seq_len, _) = inputs_embeds.dims3()?;
// let position_ids = match position_ids { let attention_mask: Option<Tensor> = {
// 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> = {
if seq_len <= 1 { if seq_len <= 1 {
None None
} else { } else {
Some(&prepare_causal_attention_mask( Some(prepare_causal_attention_mask(
b_size, b_size,
seq_len, seq_len,
0, 0,
@@ -514,9 +505,9 @@ impl HunYuanVLTextModel {
{ {
let (cos, sin) = let (cos, sin) =
get_xd_cos_sin(&cos, &sin, position_ids, self.xdrope_section.clone())?; 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 { } 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)?; let xs = self.norm.forward(&xs)?;
+1 -1
View File
@@ -119,7 +119,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
decode_ids.extend_from_slice(&error_tokens); decode_ids.extend_from_slice(&error_tokens);
} }
decode_ids.push(next_token); 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("") { if decoded_token.contains("") {
error_tokens.push(next_token); error_tokens.push(next_token);
if error_tokens.len() > 3 { if error_tokens.len() > 3 {
+12 -7
View File
@@ -232,11 +232,11 @@ impl MiniCPMModel {
.embed_tokens .embed_tokens
.forward(input_ids)? .forward(input_ids)?
.affine(self.cfg.scale_emb, 0.0)?; .affine(self.cfg.scale_emb, 0.0)?;
let attention_mask: Option<&Tensor> = { let attention_mask: Option<Tensor> = {
if seq_len <= 1 { if seq_len <= 1 {
None None
} else { } else {
Some(&prepare_causal_attention_mask( Some(prepare_causal_attention_mask(
bs, bs,
seq_len, seq_len,
0, 0,
@@ -248,7 +248,8 @@ impl MiniCPMModel {
let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?;
let mut hidden_states = input_embeds; let mut hidden_states = input_embeds;
for decode_layer in &self.layers { 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)?; hidden_states = self.norm.forward(&hidden_states)?;
let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?; let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?;
@@ -266,11 +267,11 @@ impl MiniCPMModel {
.embed_tokens .embed_tokens
.forward(input_ids)? .forward(input_ids)?
.affine(self.cfg.scale_emb, 0.0)?; .affine(self.cfg.scale_emb, 0.0)?;
let attention_mask: Option<&Tensor> = { let attention_mask: Option<Tensor> = {
if seq_len <= 1 { if seq_len <= 1 {
None None
} else { } else {
Some(&prepare_causal_attention_mask( Some(prepare_causal_attention_mask(
bs, bs,
seq_len, seq_len,
0, 0,
@@ -281,8 +282,12 @@ impl MiniCPMModel {
let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?;
let mut hidden_states = input_embeds; let mut hidden_states = input_embeds;
for decode_layer in &mut self.layers { for decode_layer in &mut self.layers {
hidden_states = hidden_states = decode_layer.forward_with_cache(
decode_layer.forward_with_cache(&hidden_states, &cos, &sin, attention_mask)?; &hidden_states,
&cos,
&sin,
attention_mask.as_ref(),
)?;
} }
hidden_states = self.norm.forward(&hidden_states)?; hidden_states = self.norm.forward(&hidden_states)?;
let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?; let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?;
+1 -1
View File
@@ -162,7 +162,7 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
decode_ids.extend_from_slice(&error_tokens); decode_ids.extend_from_slice(&error_tokens);
} }
decode_ids.push(next_token); 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("") { if decoded_token.contains("") {
error_tokens.push(next_token); error_tokens.push(next_token);
if error_tokens.len() > 3 { if error_tokens.len() > 3 {
+4 -4
View File
@@ -376,11 +376,11 @@ impl Ernie4_5Model {
self.rope_scaling.mrope_section.clone(), self.rope_scaling.mrope_section.clone(),
)?; )?;
let mut xs = inputs_embeds.clone(); let mut xs = inputs_embeds.clone();
let attention_mask: Option<&Tensor> = { let attention_mask: Option<Tensor> = {
if seq_len <= 1 { if seq_len <= 1 {
None None
} else { } else {
Some(&prepare_causal_attention_mask( Some(prepare_causal_attention_mask(
b_size, b_size,
seq_len, seq_len,
0, 0,
@@ -389,7 +389,7 @@ impl Ernie4_5Model {
} }
}; };
for layer in self.layers.iter_mut() { 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)?; let xs = xs.apply(&self.norm)?;
Ok(xs) Ok(xs)
@@ -564,7 +564,7 @@ impl PaddleOCRVLModel {
} }
} }
Err(e) => { Err(e) => {
println!("get vision_indices err: {}", e); println!("get vision_indices err: {e}");
} }
}; };
+1 -1
View File
@@ -189,7 +189,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
decode_ids.extend_from_slice(&error_tokens); decode_ids.extend_from_slice(&error_tokens);
} }
decode_ids.push(next_token); 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("") { if decoded_token.contains("") {
error_tokens.push(next_token); error_tokens.push(next_token);
if error_tokens.len() > 3 { if error_tokens.len() > 3 {
+4 -4
View File
@@ -713,11 +713,11 @@ impl Qwen2_5VLTextModel {
self.rope_scaling.mrope_section.clone(), self.rope_scaling.mrope_section.clone(),
)?; )?;
let mut xs = inputs_embeds.clone(); let mut xs = inputs_embeds.clone();
let attention_mask: Option<&Tensor> = { let attention_mask: Option<Tensor> = {
if seq_len <= 1 { if seq_len <= 1 {
None None
} else { } else {
Some(&prepare_causal_attention_mask( Some(prepare_causal_attention_mask(
b_size, b_size,
seq_len, seq_len,
0, 0,
@@ -726,7 +726,7 @@ impl Qwen2_5VLTextModel {
} }
}; };
for layer in self.layers.iter_mut() { 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)?; let xs = xs.apply(&self.norm)?;
Ok(xs) Ok(xs)
@@ -898,7 +898,7 @@ impl Qwen2_5VLModel {
} }
} }
Err(e) => { Err(e) => {
println!("get vision_indices err: {}", e); println!("get vision_indices err: {e}");
} }
}; };
+2 -2
View File
@@ -243,7 +243,7 @@ impl Qwen2_5VLProcessor {
let image = get_image(file); let image = get_image(file);
match image { match image {
Ok(img) => file_vec.push(img), 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() { if !file_vec.is_empty() {
@@ -253,7 +253,7 @@ impl Qwen2_5VLProcessor {
pixel_values = Some(img_input.data); pixel_values = Some(img_input.data);
image_grid_thw = Some(img_input.grid_thw); image_grid_thw = Some(img_input.grid_thw);
} }
Err(e) => println!("img process_images err: {:?}", e), Err(e) => println!("img process_images err: {e:?}"),
}; };
} }
} }
+1 -1
View File
@@ -191,7 +191,7 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
decode_ids.extend_from_slice(&error_tokens); decode_ids.extend_from_slice(&error_tokens);
} }
decode_ids.push(next_token); 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("") { if decoded_token.contains("") {
error_tokens.push(next_token); error_tokens.push(next_token);
if error_tokens.len() > 3 { if error_tokens.len() > 3 {
+10 -6
View File
@@ -747,11 +747,11 @@ impl Qwen3VLTextModel {
self.mrope_section.clone(), self.mrope_section.clone(),
)?; )?;
let mut xs = inputs_embeds.clone(); let mut xs = inputs_embeds.clone();
let attention_mask: Option<&Tensor> = { let attention_mask: Option<Tensor> = {
if seq_len <= 1 { if seq_len <= 1 {
None None
} else { } else {
Some(&prepare_causal_attention_mask( Some(prepare_causal_attention_mask(
b_size, b_size,
seq_len, seq_len,
0, 0,
@@ -760,7 +760,7 @@ impl Qwen3VLTextModel {
} }
}; };
for (layer_idx, layer) in self.layers.iter_mut().enumerate() { 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() if let Some(deepstack_embeds) = deepstack_visual_embeds.as_ref()
&& layer_idx < deepstack_embeds.len() && layer_idx < deepstack_embeds.len()
{ {
@@ -867,7 +867,7 @@ impl Qwen3VLModel {
.reshape((*t as usize, ()))?, .reshape((*t as usize, ()))?,
); );
} }
Some(&Tensor::cat(&v_thw_vec, 0)?) Some(Tensor::cat(&v_thw_vec, 0)?)
} }
None => None, None => None,
}; };
@@ -918,7 +918,11 @@ impl Qwen3VLModel {
text_end = vision_indices_vec[j]; text_end = vision_indices_vec[j];
} }
if token == video_token_id as u32 { if token == video_token_id as u32 {
thw = video_grid_thw.unwrap().i(video_index)?.to_vec1::<u32>()?; thw = video_grid_thw
.as_ref()
.unwrap()
.i(video_index)?
.to_vec1::<u32>()?;
text_end = vision_indices_vec[j]; text_end = vision_indices_vec[j];
video_index += 1; video_index += 1;
} }
@@ -987,7 +991,7 @@ impl Qwen3VLModel {
} }
} }
Err(e) => { Err(e) => {
println!("get vision_indices err: {}", e); println!("get vision_indices err: {e}");
} }
}; };
if text_start < input_ids_i.dim(0)? as u32 { if text_start < input_ids_i.dim(0)? as u32 {
+4 -6
View File
@@ -310,7 +310,7 @@ impl Qwen3VLProcessor {
let image = get_image(file); let image = get_image(file);
match image { match image {
Ok(img) => file_vec.push(img), 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() { if !file_vec.is_empty() {
@@ -320,7 +320,7 @@ impl Qwen3VLProcessor {
pixel_values = Some(img_input.data); pixel_values = Some(img_input.data);
image_grid_thw = Some(img_input.grid_thw); 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)> { ) -> Result<(u32, u32)> {
if num_frames < temporal_factor { if num_frames < temporal_factor {
return Err(anyhow!(format!( return Err(anyhow!(format!(
"{} must be larger than temporal_factor {}", "{num_frames} must be larger than temporal_factor {temporal_factor}"
num_frames, temporal_factor
))); )));
} }
if height < factor || width < factor { if height < factor || width < factor {
return Err(anyhow!(format!( return Err(anyhow!(format!(
"height:{} or width:{} must be larger than factor:{}", "height:{height} or width:{width} must be larger than factor:{factor}"
height, width, factor
))); )));
} }
if std::cmp::max(height, width) / std::cmp::min(height, width) > 200 { if std::cmp::max(height, width) / std::cmp::min(height, width) > 200 {
+12 -7
View File
@@ -266,11 +266,11 @@ impl MiniCPMModel {
is_causal: bool, is_causal: bool,
) -> Result<Tensor> { ) -> Result<Tensor> {
let (bs, seq_len, _) = input_embeds.dims3()?; let (bs, seq_len, _) = input_embeds.dims3()?;
let attention_mask: Option<&Tensor> = { let attention_mask: Option<Tensor> = {
if !is_causal || seq_len <= 1 { if !is_causal || seq_len <= 1 {
None None
} else { } else {
Some(&prepare_causal_attention_mask( Some(prepare_causal_attention_mask(
bs, bs,
seq_len, seq_len,
0, 0,
@@ -281,7 +281,8 @@ impl MiniCPMModel {
let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?;
let mut hidden_states = input_embeds.clone(); let mut hidden_states = input_embeds.clone();
for decode_layer in &self.layers { 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)?; hidden_states = self.norm.forward(&hidden_states)?;
Ok(hidden_states) Ok(hidden_states)
@@ -298,11 +299,11 @@ impl MiniCPMModel {
_ => return Err(anyhow!("MiniCPMModelinput_embeds illigal")), _ => return Err(anyhow!("MiniCPMModelinput_embeds illigal")),
}; };
let (bs, seq_len, _) = input_embeds.dims3()?; let (bs, seq_len, _) = input_embeds.dims3()?;
let attention_mask: Option<&Tensor> = { let attention_mask: Option<Tensor> = {
if seq_len <= 1 { if seq_len <= 1 {
None None
} else { } else {
Some(&prepare_causal_attention_mask( Some(prepare_causal_attention_mask(
bs, bs,
seq_len, seq_len,
0, 0,
@@ -313,8 +314,12 @@ impl MiniCPMModel {
let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?;
let mut hidden_states = input_embeds.clone(); let mut hidden_states = input_embeds.clone();
for decode_layer in &mut self.layers { for decode_layer in &mut self.layers {
hidden_states = hidden_states = decode_layer.forward_with_cache(
decode_layer.forward_with_cache(&hidden_states, &cos, &sin, attention_mask)?; &hidden_states,
&cos,
&sin,
attention_mask.as_ref(),
)?;
} }
hidden_states = self.norm.forward(&hidden_states)?; hidden_states = self.norm.forward(&hidden_states)?;
+2 -2
View File
@@ -19,7 +19,7 @@ impl SingleChineseTokenizer {
"tokenizer.json not exists in model path" "tokenizer.json not exists in model path"
); );
let tokenizer = Tokenizer::from_file(tokenizer_file) 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(); let mut multichar_tokens = Vec::new();
for (token, _) in tokenizer.get_vocab(false) { for (token, _) in tokenizer.get_vocab(false) {
let len = token.chars().count(); let len = token.chars().count();
@@ -42,7 +42,7 @@ impl SingleChineseTokenizer {
let encode = self let encode = self
.tokenizer .tokenizer
.encode(text, false) .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(); let tokens = encode.get_tokens();
// println!("tokens: {:?}", tokens); // println!("tokens: {:?}", tokens);
let mut split_character = Vec::new(); let mut split_character = Vec::new();