generate remove unnecessary code
This commit is contained in:
@@ -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)?;
|
||||||
|
|||||||
@@ -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)?;
|
||||||
|
|||||||
@@ -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)?;
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user