diff --git a/src/models/qwen3vl/model.rs b/src/models/qwen3vl/model.rs index 518a8b5..e556fb6 100644 --- a/src/models/qwen3vl/model.rs +++ b/src/models/qwen3vl/model.rs @@ -516,19 +516,11 @@ impl Qwen3VLVisionModel { let sin = emb.sin()?; let cu_seqlens = grid_thw.i((.., 1))?.mul(&grid_thw.i((.., 2))?)?; let grid_t = grid_thw.i((.., 0))?.to_vec1::()?; - let cu_seqlens_full = match cu_seqlens.rank() { - 1 => cu_seqlens.repeat(grid_t[0] as usize)?, - 2 => { - let mut cu_seqlens_repeat = Vec::new(); - for (index, t) in grid_t.iter().enumerate() { - cu_seqlens_repeat.push(cu_seqlens.i(index)?.repeat(*t as usize)?); - } - Tensor::cat(&cu_seqlens_repeat, 0)?.flatten_all()? - } - _ => { - return Err(anyhow!(format!("create cu_seqlens error"))); - } - }; + let mut cu_seqlens_repeat = Vec::new(); + for (index, t) in grid_t.iter().enumerate() { + cu_seqlens_repeat.push(cu_seqlens.i(index)?.repeat(*t as usize)?); + } + let cu_seqlens_full = Tensor::cat(&cu_seqlens_repeat, 0)?.flatten_all()?; let cu_seqlens = cu_seqlens_full .to_dtype(DType::F64)? .cumsum(0)? @@ -895,6 +887,8 @@ impl Qwen3VLModel { let mut v_thw_vec = Vec::new(); for (index, t) in grid_t.iter().enumerate() { let mut thw_i = thw.i(index)?.to_vec1::()?; + // [12, 30, 50] + // [1, 30, 50]*t thw_i[0] = 1; v_thw_vec.push( Tensor::new(thw_i, thw.device())? diff --git a/tests/messy_test.rs b/tests/messy_test.rs index 86fec2c..a64b3c8 100644 --- a/tests/messy_test.rs +++ b/tests/messy_test.rs @@ -1,14 +1,20 @@ use aha::utils::tensor_utils::bitor_tensor; use anyhow::Result; -use candle_core::Tensor; +use candle_core::{IndexOp, Tensor}; #[test] fn messy_test() -> Result<()> { let device = &candle_core::Device::Cpu; - let image_mask = Tensor::new(vec![0u32, 0, 0, 1, 0, 1], device)?; - let video_mask = Tensor::new(vec![0u32, 1, 0, 1, 0, 1], device)?; - let visual_mask = bitor_tensor(&image_mask, &video_mask)?; - println!("visual_mask: {}", visual_mask); + let grid_thw = Tensor::new(vec![vec![3u32, 12, 20], vec![5, 30, 25]], device)?; + let cu_seqlens = grid_thw.i((.., 1))?.mul(&grid_thw.i((.., 2))?)?; + let grid_t = grid_thw.i((.., 0))?.to_vec1::()?; + println!("cu_seqlens: {}", cu_seqlens); + println!("cu_seqlens rank: {}", cu_seqlens.rank()); + println!("grid_t: {:?}", grid_t); + // let image_mask = Tensor::new(vec![0u32, 0, 0, 1, 0, 1], device)?; + // let video_mask = Tensor::new(vec![0u32, 1, 0, 1, 0, 1], device)?; + // let visual_mask = bitor_tensor(&image_mask, &video_mask)?; + // println!("visual_mask: {}", visual_mask); // let x = Tensor::arange_step(0.0_f32, 5., 0.5, &device)?; // let x_int = x.to_dtype(candle_core::DType::U32)?; // println!("x: {}", x);