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
+19 -13
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
.processor
.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 seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut generate = Vec::new();
let logits = self.deepseekocr_model.forward(
&input_ids,
Some(&images_ori),
Some(&image_crop),
Some(&images_seq_mask),
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 0..sample_len {
for _ in 1..sample_len {
let logits = self.deepseekocr_model.forward(
&input_ids,
images_ori,
image_crop,
images_seq_mask,
images_spatial_crop_t,
None,
None,
None,
None,
seqlen_offset,
)?;
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
@@ -107,10 +117,6 @@ impl GenerateModel for DeepseekOCRGenerateModel {
seqlen_offset += seq_len;
seq_len = 1;
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 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)?;
let mut seq_len = input_ids.dim(1)?;
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 cache_position = Tensor::ones_like(&input_ids.i(0)?)?
.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 stream = stream! {
let mut error_tokens = Vec::new();
let mut pixel_values = pixel_values.as_ref();
let image_grid_thw = image_grid_thw.as_ref();
let mut pixel_values_video = pixel_values_video.as_ref();
let video_grid_thw = video_grid_thw.as_ref();
let mut pixel_values = input.pixel_values.as_ref();
let image_grid_thw = input.image_grid_thw.as_ref();
let mut pixel_values_video = input.pixel_values_video.as_ref();
let video_grid_thw = input.video_grid_thw.as_ref();
let mut tool_call_id = None;
let mut tool_call_content = String::new();
for _ in 0..sample_len {
@@ -176,7 +170,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
&mask,
Some(&cache_position),
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 next_token = logit_processor.sample(&logits)?;
-1
View File
@@ -116,7 +116,6 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
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
.chat_template
.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)?;
let mut seq_len = input_ids.dim(1)?;
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 sample_len = mes.max_tokens.unwrap_or(1024);
let stream = stream! {
let mut error_tokens = Vec::new();
let mut pixel_values = pixel_values.as_ref();
let image_grid_thw = image_grid_thw.as_ref();
let mut pixel_values_video = pixel_values_video.as_ref();
let video_grid_thw = video_grid_thw.as_ref();
let mut pixel_values = input.pixel_values.as_ref();
let image_grid_thw = input.image_grid_thw.as_ref();
let mut pixel_values_video = input.pixel_values_video.as_ref();
let video_grid_thw = input.video_grid_thw.as_ref();
let mut tool_call_id = None;
let mut tool_call_content = String::new();
for _ in 0..sample_len {
+1 -1
View File
@@ -63,7 +63,7 @@ fn qwen2_5vl_generate() -> Result<()> {
#[tokio::test]
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 dtype = DType::BF16;