generate remove unnecessary code

This commit is contained in:
jhqxxx
2026-03-06 19:40:39 +08:00
parent 037207eccd
commit 2ef86b5a88
5 changed files with 29 additions and 34 deletions
+20 -14
View File
@@ -80,22 +80,32 @@ impl GenerateModel for DeepseekOCRGenerateModel {
let (mut input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self let (mut input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self
.processor .processor
.process_info(&mes, &self.tokenizer, base_size, image_size, crop_mode)?; .process_info(&mes, &self.tokenizer, base_size, image_size, crop_mode)?;
let mut images_ori = Some(&images_ori);
let mut image_crop = Some(&image_crop);
let mut images_seq_mask = Some(&images_seq_mask);
let mut images_spatial_crop_t = Some(&images_spatial_crop_t);
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32; let prompt_tokens = seq_len as u32;
let mut generate = Vec::new(); let mut generate = Vec::new();
let sample_len = mes.max_tokens.unwrap_or(1024);
for _ in 0..sample_len {
let logits = self.deepseekocr_model.forward( let logits = self.deepseekocr_model.forward(
&input_ids, &input_ids,
images_ori, Some(&images_ori),
image_crop, Some(&image_crop),
images_seq_mask, Some(&images_seq_mask),
images_spatial_crop_t, Some(&images_spatial_crop_t),
seqlen_offset,
)?;
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
let next_token = logit_processor.sample(&logits)?;
generate.push(next_token);
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
seqlen_offset += seq_len;
seq_len = 1;
let sample_len = mes.max_tokens.unwrap_or(1024);
for _ in 1..sample_len {
let logits = self.deepseekocr_model.forward(
&input_ids,
None,
None,
None,
None,
seqlen_offset, seqlen_offset,
)?; )?;
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
@@ -107,10 +117,6 @@ impl GenerateModel for DeepseekOCRGenerateModel {
seqlen_offset += seq_len; seqlen_offset += seq_len;
seq_len = 1; seq_len = 1;
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
images_ori = None;
image_crop = None;
images_seq_mask = None;
images_spatial_crop_t = None;
} }
let num_token = generate.len() as u32; let num_token = generate.len() as u32;
let res = self.tokenizer.token_decode(generate)?; let res = self.tokenizer.token_decode(generate)?;
+5 -11
View File
@@ -144,12 +144,6 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
.text_encode(input.replace_text.clone(), &self.device)?; .text_encode(input.replace_text.clone(), &self.device)?;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let pixel_values = input.pixel_values.clone();
let image_grid_thw = input.image_grid_thw.clone();
let pixel_values_video = input.pixel_values_video.clone();
let video_grid_thw = input.video_grid_thw.clone();
let second_per_grid_ts = input.second_per_grid_ts.clone();
let mut mask = Tensor::ones_like(&input_ids)?; let mut mask = Tensor::ones_like(&input_ids)?;
let mut cache_position = Tensor::ones_like(&input_ids.i(0)?)? let mut cache_position = Tensor::ones_like(&input_ids.i(0)?)?
.to_dtype(candle_core::DType::F64)? .to_dtype(candle_core::DType::F64)?
@@ -160,10 +154,10 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
let sample_len = mes.max_tokens.unwrap_or(512); let sample_len = mes.max_tokens.unwrap_or(512);
let stream = stream! { let stream = stream! {
let mut error_tokens = Vec::new(); let mut error_tokens = Vec::new();
let mut pixel_values = pixel_values.as_ref(); let mut pixel_values = input.pixel_values.as_ref();
let image_grid_thw = image_grid_thw.as_ref(); let image_grid_thw = input.image_grid_thw.as_ref();
let mut pixel_values_video = pixel_values_video.as_ref(); let mut pixel_values_video = input.pixel_values_video.as_ref();
let video_grid_thw = video_grid_thw.as_ref(); let video_grid_thw = input.video_grid_thw.as_ref();
let mut tool_call_id = None; let mut tool_call_id = None;
let mut tool_call_content = String::new(); let mut tool_call_content = String::new();
for _ in 0..sample_len { for _ in 0..sample_len {
@@ -176,7 +170,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
&mask, &mask,
Some(&cache_position), Some(&cache_position),
seqlen_offset, seqlen_offset,
second_per_grid_ts.clone(), input.second_per_grid_ts.clone(),
)?; )?;
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
let next_token = logit_processor.sample(&logits)?; let next_token = logit_processor.sample(&logits)?;
-1
View File
@@ -116,7 +116,6 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
let mut logit_processor = let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking"); let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
// let mes_render = self.chat_template.apply_chat_template(&mes)?;
let mes_render = self let mes_render = self
.chat_template .chat_template
.apply_chat_temp_think(&mes, enable_thinking)?; .apply_chat_temp_think(&mes, enable_thinking)?;
+4 -8
View File
@@ -153,18 +153,14 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
.text_encode(input.replace_text.clone(), &self.device)?; .text_encode(input.replace_text.clone(), &self.device)?;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
let pixel_values = input.pixel_values.clone();
let image_grid_thw = input.image_grid_thw.clone();
let pixel_values_video = input.pixel_values_video.clone();
let video_grid_thw = input.video_grid_thw.clone();
let mut cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?; let mut cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?;
let sample_len = mes.max_tokens.unwrap_or(1024); let sample_len = mes.max_tokens.unwrap_or(1024);
let stream = stream! { let stream = stream! {
let mut error_tokens = Vec::new(); let mut error_tokens = Vec::new();
let mut pixel_values = pixel_values.as_ref(); let mut pixel_values = input.pixel_values.as_ref();
let image_grid_thw = image_grid_thw.as_ref(); let image_grid_thw = input.image_grid_thw.as_ref();
let mut pixel_values_video = pixel_values_video.as_ref(); let mut pixel_values_video = input.pixel_values_video.as_ref();
let video_grid_thw = video_grid_thw.as_ref(); let video_grid_thw = input.video_grid_thw.as_ref();
let mut tool_call_id = None; let mut tool_call_id = None;
let mut tool_call_content = String::new(); let mut tool_call_content = String::new();
for _ in 0..sample_len { for _ in 0..sample_len {
+1 -1
View File
@@ -63,7 +63,7 @@ fn qwen2_5vl_generate() -> Result<()> {
#[tokio::test] #[tokio::test]
async fn qwen2_5vl_stream() -> Result<()> { async fn qwen2_5vl_stream() -> Result<()> {
// test with cuda+flash-attn: RUST_BACKTRACE=1 cargo test -F cuda,flash-attn qwen2_5vl_stream -r -- --nocapture // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda qwen2_5vl_stream -r -- --nocapture
// let device = Device::cuda_if_available(0)?; // let device = Device::cuda_if_available(0)?;
// let dtype = DType::BF16; // let dtype = DType::BF16;