update cu_seqlens_repeat
This commit is contained in:
@@ -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
@@ -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);
|
||||||
|
|||||||
Reference in New Issue
Block a user