refactor: remove unstable Option<&T> branch by returning owned Tensor
This commit is contained in:
@@ -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 {
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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)?;
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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)?;
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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}");
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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}");
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -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:?}"),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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)?;
|
||||||
|
|
||||||
|
|||||||
@@ -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();
|
||||||
|
|||||||
Reference in New Issue
Block a user