update video with image

This commit is contained in:
jhqxxx
2025-10-26 23:19:11 +08:00
parent 0e4ca66aca
commit 86e814708a
3 changed files with 58 additions and 45 deletions
+17 -13
View File
@@ -415,11 +415,12 @@ impl Qwen3VLVisionModel {
let mut patch_pos_embeds_permute = vec![]; let mut patch_pos_embeds_permute = vec![];
let patch_pos_embeds = split_tensor(&patch_pos_embeds, &split_idx, 0)?; let patch_pos_embeds = split_tensor(&patch_pos_embeds, &split_idx, 0)?;
let merge_size = self.spatial_merge_size; 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::<u32>()?[..] else { let [t, h, w] = grid_thw.i(i)?.to_vec1::<u32>()?[..] else {
return Err(anyhow!(format!("grid_thw Expected exactly 3 elements"))); 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_emebd_last_dim = pos_embed.dim(D::Minus1)?;
let pos_embed = pos_embed.repeat((t as usize, 1))?; let pos_embed = pos_embed.repeat((t as usize, 1))?;
let shape = Shape::from(vec![ let shape = Shape::from(vec![
@@ -803,13 +804,13 @@ 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)?;
if deepstack_visual_embeds.is_some() if let Some(deepstack_embeds) = deepstack_visual_embeds.as_ref()
&& layer_idx < deepstack_visual_embeds.as_ref().unwrap().len() && layer_idx < deepstack_embeds.len()
{ {
xs = mask_index_add( xs = mask_index_add(
&xs.squeeze(0)?, &xs.squeeze(0)?,
&visual_pos_masks.unwrap().squeeze(0)?, &visual_pos_masks.unwrap().squeeze(0)?,
&deepstack_visual_embeds.as_ref().unwrap()[layer_idx], &deepstack_embeds[layer_idx],
)? )?
.unsqueeze(0)?; .unsqueeze(0)?;
} }
@@ -1169,9 +1170,11 @@ impl Qwen3VLModel {
} }
let mut visual_pos_mask = None; let mut visual_pos_mask = None;
let mut deepstack_visual_embeds = None; let mut deepstack_visual_embeds = None;
if image_mask.is_some() && video_mask.is_some() { // if image_mask.is_some() && video_mask.is_some() {
let image_mask_ = image_mask.unwrap(); if let Some(image_mask_) = image_mask {
let video_mask_ = video_mask.unwrap(); 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_mask = bitor_tensor(&image_mask_, &video_mask_)?;
let visual_none_zero_index = nonzero_index(&visual_mask)?; let visual_none_zero_index = nonzero_index(&visual_mask)?;
let image_mask_joint = image_mask_.gather(&visual_none_zero_index, 0)?; let image_mask_joint = image_mask_.gather(&visual_none_zero_index, 0)?;
@@ -1194,13 +1197,14 @@ impl Qwen3VLModel {
let embed_joint = embed_joint.index_add(&video_nonzero_joint, vid_embed, 0)?; let embed_joint = embed_joint.index_add(&video_nonzero_joint, vid_embed, 0)?;
deepstack_embeds.push(embed_joint); deepstack_embeds.push(embed_joint);
} }
visual_pos_mask = Some(visual_mask); visual_pos_mask = Some(visual_mask.unsqueeze(0)?);
deepstack_visual_embeds = Some(deepstack_embeds); deepstack_visual_embeds = Some(deepstack_embeds);
} else if image_mask.is_some() { } else {
visual_pos_mask = image_mask; visual_pos_mask = Some(image_mask_);
deepstack_visual_embeds = deepstack_image_embeds; 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; deepstack_visual_embeds = deepstack_video_embeds;
} }
+4 -2
View File
@@ -217,8 +217,10 @@ impl Qwen3VLTextRotaryEmbedding {
) -> Result<Tensor> { ) -> Result<Tensor> {
let mut freqs_t = freqs.i(0)?.contiguous()?; //(3, bs, seq_len, head_dim //2) -> (bs, seq_len, head_dim //2) 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 { // for dim in 1..3 {
let length = mrope_section[dim] * 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 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 = freqs.i(dim)?.contiguous()?; // (bs, seq_len, head_dim //2)
let src = src.index_select(&idx, D::Minus1)?.contiguous()?; let src = src.index_select(&idx, D::Minus1)?.contiguous()?;
+12 -5
View File
@@ -25,9 +25,16 @@ fn qwen3vl_generate() -> Result<()> {
"url": "./assets/video/video_test.mp4" "url": "./assets/video/video_test.mp4"
} }
}, },
{
"type": "image",
"image_url":
{
"url": "file://./assets/img/voxcpm.png"
}
},
{ {
"type": "text", "type": "text",
"text": "视频中发生了什么." "text": "描述视频和图片内容"
} }
] ]
} }
@@ -62,15 +69,15 @@ async fn qwen3vl_stream() -> Result<()> {
"role": "user", "role": "user",
"content": [ "content": [
{ {
"type": "image", "type": "video",
"image_url": "video_url":
{ {
"url": "file://./assets/img/voxcpm.png" "url": "./assets/video/video_test.mp4"
} }
}, },
{ {
"type": "text", "type": "text",
"text": "描述这张图片" "text": "视频中发生了什么?"
} }
] ]
} }