update cu_seqlens_repeat

This commit is contained in:
jhqxxx
2025-11-04 12:01:41 +08:00
parent b42c5c718c
commit 082bf544d8
2 changed files with 18 additions and 18 deletions
+7 -13
View File
@@ -516,19 +516,11 @@ impl Qwen3VLVisionModel {
let sin = emb.sin()?; let sin = emb.sin()?;
let cu_seqlens = grid_thw.i((.., 1))?.mul(&grid_thw.i((.., 2))?)?; let cu_seqlens = grid_thw.i((.., 1))?.mul(&grid_thw.i((.., 2))?)?;
let grid_t = grid_thw.i((.., 0))?.to_vec1::<u32>()?; let grid_t = grid_thw.i((.., 0))?.to_vec1::<u32>()?;
let cu_seqlens_full = match cu_seqlens.rank() { let mut cu_seqlens_repeat = Vec::new();
1 => cu_seqlens.repeat(grid_t[0] as usize)?, for (index, t) in grid_t.iter().enumerate() {
2 => { cu_seqlens_repeat.push(cu_seqlens.i(index)?.repeat(*t as usize)?);
let mut cu_seqlens_repeat = Vec::new(); }
for (index, t) in grid_t.iter().enumerate() { let cu_seqlens_full = Tensor::cat(&cu_seqlens_repeat, 0)?.flatten_all()?;
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 cu_seqlens = cu_seqlens_full let cu_seqlens = cu_seqlens_full
.to_dtype(DType::F64)? .to_dtype(DType::F64)?
.cumsum(0)? .cumsum(0)?
@@ -895,6 +887,8 @@ impl Qwen3VLModel {
let mut v_thw_vec = Vec::new(); let mut v_thw_vec = Vec::new();
for (index, t) in grid_t.iter().enumerate() { for (index, t) in grid_t.iter().enumerate() {
let mut thw_i = thw.i(index)?.to_vec1::<u32>()?; let mut thw_i = thw.i(index)?.to_vec1::<u32>()?;
// [12, 30, 50]
// [1, 30, 50]*t
thw_i[0] = 1; thw_i[0] = 1;
v_thw_vec.push( v_thw_vec.push(
Tensor::new(thw_i, thw.device())? Tensor::new(thw_i, thw.device())?
+11 -5
View File
@@ -1,14 +1,20 @@
use aha::utils::tensor_utils::bitor_tensor; use aha::utils::tensor_utils::bitor_tensor;
use anyhow::Result; use anyhow::Result;
use candle_core::Tensor; use candle_core::{IndexOp, Tensor};
#[test] #[test]
fn messy_test() -> Result<()> { fn messy_test() -> Result<()> {
let device = &candle_core::Device::Cpu; let device = &candle_core::Device::Cpu;
let image_mask = Tensor::new(vec![0u32, 0, 0, 1, 0, 1], device)?; let grid_thw = Tensor::new(vec![vec![3u32, 12, 20], vec![5, 30, 25]], device)?;
let video_mask = Tensor::new(vec![0u32, 1, 0, 1, 0, 1], device)?; let cu_seqlens = grid_thw.i((.., 1))?.mul(&grid_thw.i((.., 2))?)?;
let visual_mask = bitor_tensor(&image_mask, &video_mask)?; let grid_t = grid_thw.i((.., 0))?.to_vec1::<u32>()?;
println!("visual_mask: {}", visual_mask); 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 = Tensor::arange_step(0.0_f32, 5., 0.5, &device)?;
// let x_int = x.to_dtype(candle_core::DType::U32)?; // let x_int = x.to_dtype(candle_core::DType::U32)?;
// println!("x: {}", x); // println!("x: {}", x);