From 86e814708a49b2a32657e46887dd5cc7f0ddefcf Mon Sep 17 00:00:00 2001 From: jhqxxx Date: Sun, 26 Oct 2025 23:19:11 +0800 Subject: [PATCH] update video with image --- src/models/qwen3vl/model.rs | 76 +++++++++++++++++++------------------ src/position_embed/rope.rs | 6 ++- tests/test_qwen3vl.rs | 21 ++++++---- 3 files changed, 58 insertions(+), 45 deletions(-) diff --git a/src/models/qwen3vl/model.rs b/src/models/qwen3vl/model.rs index 5b2c8f6..e826efc 100644 --- a/src/models/qwen3vl/model.rs +++ b/src/models/qwen3vl/model.rs @@ -415,11 +415,12 @@ impl Qwen3VLVisionModel { let mut patch_pos_embeds_permute = vec![]; let patch_pos_embeds = split_tensor(&patch_pos_embeds, &split_idx, 0)?; let merge_size = self.spatial_merge_size; - for i in 0..grid_thw.dim(0)? { + // for i in 0..grid_thw.dim(0)? { + for (i, pos_embed) in patch_pos_embeds.iter().enumerate() { let [t, h, w] = grid_thw.i(i)?.to_vec1::()?[..] else { return Err(anyhow!(format!("grid_thw Expected exactly 3 elements"))); }; - let pos_embed = &patch_pos_embeds[i]; + // let pos_embed = &patch_pos_embeds[i]; let pos_emebd_last_dim = pos_embed.dim(D::Minus1)?; let pos_embed = pos_embed.repeat((t as usize, 1))?; let shape = Shape::from(vec![ @@ -803,13 +804,13 @@ impl Qwen3VLTextModel { }; for (layer_idx, layer) in self.layers.iter_mut().enumerate() { xs = layer.forward(&xs, &cos, &sin, attention_mask)?; - if deepstack_visual_embeds.is_some() - && layer_idx < deepstack_visual_embeds.as_ref().unwrap().len() + if let Some(deepstack_embeds) = deepstack_visual_embeds.as_ref() + && layer_idx < deepstack_embeds.len() { xs = mask_index_add( &xs.squeeze(0)?, &visual_pos_masks.unwrap().squeeze(0)?, - &deepstack_visual_embeds.as_ref().unwrap()[layer_idx], + &deepstack_embeds[layer_idx], )? .unsqueeze(0)?; } @@ -1169,38 +1170,41 @@ impl Qwen3VLModel { } let mut visual_pos_mask = None; let mut deepstack_visual_embeds = None; - if image_mask.is_some() && video_mask.is_some() { - let image_mask_ = image_mask.unwrap(); - let video_mask_ = video_mask.unwrap(); - let visual_mask = bitor_tensor(&image_mask_, &video_mask_)?; - let visual_none_zero_index = nonzero_index(&visual_mask)?; - let image_mask_joint = image_mask_.gather(&visual_none_zero_index, 0)?; - let image_nonzero_joint = nonzero_index(&image_mask_joint)?; - let video_mask_joint = video_mask_.gather(&visual_none_zero_index, 0)?; - let video_nonzero_joint = nonzero_index(&video_mask_joint)?; - let mut deepstack_embeds = vec![]; - let visual_len = visual_none_zero_index.dim(0)?; - for (img_embed, vid_embed) in deepstack_image_embeds - .unwrap() - .iter() - .zip(deepstack_video_embeds.unwrap().iter()) - { - let embed_joint = Tensor::zeros( - (visual_len, img_embed.dim(D::Minus1)?), - img_embed.dtype(), - img_embed.device(), - )?; - let embed_joint = embed_joint.index_add(&image_nonzero_joint, img_embed, 0)?; - let embed_joint = embed_joint.index_add(&video_nonzero_joint, vid_embed, 0)?; - deepstack_embeds.push(embed_joint); + // if image_mask.is_some() && video_mask.is_some() { + if let Some(image_mask_) = image_mask { + if let Some(video_mask_) = video_mask { + let image_mask_ = image_mask_.squeeze(0)?; + let video_mask_ = video_mask_.squeeze(0)?; + let visual_mask = bitor_tensor(&image_mask_, &video_mask_)?; + let visual_none_zero_index = nonzero_index(&visual_mask)?; + let image_mask_joint = image_mask_.gather(&visual_none_zero_index, 0)?; + let image_nonzero_joint = nonzero_index(&image_mask_joint)?; + let video_mask_joint = video_mask_.gather(&visual_none_zero_index, 0)?; + let video_nonzero_joint = nonzero_index(&video_mask_joint)?; + let mut deepstack_embeds = vec![]; + let visual_len = visual_none_zero_index.dim(0)?; + for (img_embed, vid_embed) in deepstack_image_embeds + .unwrap() + .iter() + .zip(deepstack_video_embeds.unwrap().iter()) + { + let embed_joint = Tensor::zeros( + (visual_len, img_embed.dim(D::Minus1)?), + img_embed.dtype(), + img_embed.device(), + )?; + let embed_joint = embed_joint.index_add(&image_nonzero_joint, img_embed, 0)?; + let embed_joint = embed_joint.index_add(&video_nonzero_joint, vid_embed, 0)?; + deepstack_embeds.push(embed_joint); + } + visual_pos_mask = Some(visual_mask.unsqueeze(0)?); + deepstack_visual_embeds = Some(deepstack_embeds); + } else { + visual_pos_mask = Some(image_mask_); + deepstack_visual_embeds = deepstack_image_embeds; } - visual_pos_mask = Some(visual_mask); - deepstack_visual_embeds = Some(deepstack_embeds); - } else if image_mask.is_some() { - visual_pos_mask = image_mask; - deepstack_visual_embeds = deepstack_image_embeds; - } else if video_mask.is_some() { - visual_pos_mask = video_mask; + } else if let Some(video_mask_) = video_mask { + visual_pos_mask = Some(video_mask_); deepstack_visual_embeds = deepstack_video_embeds; } diff --git a/src/position_embed/rope.rs b/src/position_embed/rope.rs index 2e9f4bb..f62d22f 100644 --- a/src/position_embed/rope.rs +++ b/src/position_embed/rope.rs @@ -217,8 +217,10 @@ impl Qwen3VLTextRotaryEmbedding { ) -> Result { let mut freqs_t = freqs.i(0)?.contiguous()?; //(3, bs, seq_len, head_dim //2) -> (bs, seq_len, head_dim //2) - for dim in 1..3 { - let length = mrope_section[dim] * 3; + // for dim in 1..3 { + for (dim, section) in mrope_section.iter().enumerate().skip(1) { + // let length = mrope_section[dim] * 3; + let length = section * 3; let idx = Tensor::arange_step(dim as u32, length as u32, 3, freqs.device())?; let src = freqs.i(dim)?.contiguous()?; // (bs, seq_len, head_dim //2) let src = src.index_select(&idx, D::Minus1)?.contiguous()?; diff --git a/tests/test_qwen3vl.rs b/tests/test_qwen3vl.rs index db93581..0d383e9 100644 --- a/tests/test_qwen3vl.rs +++ b/tests/test_qwen3vl.rs @@ -17,17 +17,24 @@ fn qwen3vl_generate() -> Result<()> { "messages": [ { "role": "user", - "content": [ + "content": [ { "type": "video", "video_url": { "url": "./assets/video/video_test.mp4" } - }, + }, + { + "type": "image", + "image_url": + { + "url": "file://./assets/img/voxcpm.png" + } + }, { "type": "text", - "text": "视频中发生了什么." + "text": "描述视频和图片内容" } ] } @@ -62,15 +69,15 @@ async fn qwen3vl_stream() -> Result<()> { "role": "user", "content": [ { - "type": "image", - "image_url": + "type": "video", + "video_url": { - "url": "file://./assets/img/voxcpm.png" + "url": "./assets/video/video_test.mp4" } }, { "type": "text", - "text": "描述这张图片" + "text": "视频中发生了什么?" } ] }