diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs index f306893..155fbb9 100644 --- a/src/models/common/mod.rs +++ b/src/models/common/mod.rs @@ -316,6 +316,7 @@ pub fn eager_attention_forward( attn_output } }; + //(b, n_head, seq_len, dim) -> (b, seq_len, n_head, dim) let attn_output = attn_output.transpose(1, 2)?.contiguous()?; Ok(attn_output) diff --git a/src/models/deepseek_ocr/model.rs b/src/models/deepseek_ocr/model.rs index de002ad..10d2f5b 100644 --- a/src/models/deepseek_ocr/model.rs +++ b/src/models/deepseek_ocr/model.rs @@ -1,11 +1,14 @@ use anyhow::{Ok, Result}; -use candle_core::{IndexOp, Tensor}; +use candle_core::{D, IndexOp, Tensor}; use candle_nn::{ - Conv2d, Conv2dConfig, Init, LayerNorm, Linear, Module, VarBuilder, conv2d, linear, - linear_no_bias, + Activation, Conv2d, Conv2dConfig, Init, LayerNorm, LayerNormConfig, Linear, Module, VarBuilder, + conv2d, layer_norm, linear, linear_no_bias, }; -use crate::models::deepseek_ocr::config::DeepseekOCRConfig; +use crate::{ + models::{common::eager_attention_forward, deepseek_ocr::config::DeepseekOCRConfig}, + utils::tensor_utils::{index_select_2d, interpolate_linear}, +}; pub struct PatchEmbed { proj: Conv2d, @@ -93,43 +96,265 @@ impl Attention { }) } - // fn get_rel_pos(q_size: usize, k_size: usize, rel_pos: &Tensor) -> Result { - // let max_rel_dist = 2 * std::cmp::max(q_size, k_size) - 1; - // let rel_pos_resized = if rel_pos.dim(0)? != max_rel_dist { - // let dtype = rel_pos.dtype(); - // let rel_pos = rel_pos.to_dtype(candle_core::DType::F32)?; - // let rel_pos_resized = - // } - // } - // fn add_decomposed_rel_pos(&self, q: &Tensor, rel_pos_h: &Tensor, rel_pos_w: &Tensor, q_size: (usize, usize), k_size: (usize, usize)) -> Result { - // let (q_h, q_w) = q_size; - // let (k_h, k_w) = k_size; - - // } + fn get_rel_pos(&self, q_size: usize, k_size: usize, rel_pos: &Tensor) -> Result { + let max_rel_dist = 2 * std::cmp::max(q_size, k_size) - 1; + let rel_pos_resized = if rel_pos.dim(0)? != max_rel_dist { + let rel_pos = rel_pos + .to_dtype(candle_core::DType::F32)? + .t()? + .unsqueeze(0)? + .contiguous()?; + let rel_pos_resized = interpolate_linear(&rel_pos, max_rel_dist, None)?; + let rel_pos_resized = rel_pos_resized + .squeeze(0)? + .t()? + .contiguous()? + .to_dtype(rel_pos.dtype())?; + rel_pos_resized + } else { + rel_pos.clone() + }; + let q_coords = Tensor::arange(0 as f32, q_size as f32, rel_pos.device())? + .unsqueeze(D::Minus1)? + .affine((k_size as f64 / q_size as f64).max(1.0), 0.0)?; + let k_coords = Tensor::arange(0 as f32, k_size as f32, rel_pos.device())? + .unsqueeze(0)? + .affine((q_size as f64 / k_size as f64).max(1.0), 0.0)?; + let relative_coords = q_coords + .broadcast_sub(&k_coords)? + .affine(1.0, (k_size - 1) as f64)? + .affine((q_size as f64 / k_size as f64).max(1.0), 0.0)?; + let relative_coords = relative_coords + .to_dtype(candle_core::DType::U32)? + .contiguous()?; + let rel_pos_resized = rel_pos_resized.contiguous()?; + let res = index_select_2d(&rel_pos_resized, &relative_coords)?; + Ok(res) + } - // pub fn forward(&mut self, xs: &Tensor) -> Result { - // let (b, h, w, _) = xs.dims4()?; - // // (3, B, n_head, h*w, head_dim) - // let qkv = self - // .qkv - // .forward(xs)? - // .reshape((b, h * w, 3, self.num_heads, ()))? - // .permute((2, 0, 3, 1, 4))? - // .contiguous()?; - // let query_states = qkv.i(0)?.contiguous()?; - // let key_states = qkv.i(1)?.contiguous()?; - // let value_states = qkv.i(2)?.contiguous()?; - // let xs = if self.use_rel_pos { - // let (rel_h, rel_w) = - // } else { + fn add_decomposed_rel_pos( + &self, + q: &Tensor, + rel_pos_h: &Tensor, + rel_pos_w: &Tensor, + q_size: (usize, usize), + k_size: (usize, usize), + ) -> Result<(Tensor, Tensor)> { + let (q_h, q_w) = q_size; + let (k_h, k_w) = k_size; + let rh = self.get_rel_pos(q_h, k_h, rel_pos_h)?; // (h, w, dim) + let rh = rh.t()?; // (h, dim, w) + let rw = self.get_rel_pos(q_w, k_w, rel_pos_w)?; + let rw = rw.t()?; + let (b, _, dim) = q.dims3()?; + let r_q = q.reshape((b, q_h, q_w, dim))?; + let rel_h = r_q.broadcast_matmul(&rh)?; + let rel_w = r_q.broadcast_matmul(&rw)?; + let rel_h = rel_h + .unsqueeze(D::Minus1)? + .reshape((b, q_h * q_w, k_h, 1))?; + let rel_w = rel_w + .unsqueeze(D::Minus2)? + .reshape((b, q_h * q_w, 1, k_w))?; + Ok((rel_h, rel_w)) + } - // } - // } + pub fn forward(&mut self, xs: &Tensor) -> Result { + let (b, h, w, _) = xs.dims4()?; + // (3, B, n_head, h*w, head_dim) + let qkv = self + .qkv + .forward(xs)? + .reshape((b, h * w, 3, self.num_heads, ()))? + .permute((2, 0, 3, 1, 4))? + .contiguous()?; + let query_states = qkv.i(0)?.contiguous()?; + let key_states = qkv.i(1)?.contiguous()?; + let value_states = qkv.i(2)?.contiguous()?; + let xs = if self.use_rel_pos { + let q_reshape = query_states.reshape((b * self.num_heads, h * w, ()))?; + let (rel_h, rel_w) = self.add_decomposed_rel_pos( + &q_reshape, + self.rel_pos_h.as_ref().unwrap(), + self.rel_pos_w.as_ref().unwrap(), + (h, w), + (h, w), + )?; + let (_, rel_h_dim1, rel_h_dim2, rel_h_dim3) = rel_h.dims4()?; + let rel_h = rel_h.reshape((b, self.num_heads, rel_h_dim1, rel_h_dim2, rel_h_dim3))?; + let (_, rel_w_dim1, rel_w_dim2, rel_w_dim3) = rel_w.dims4()?; + let rel_w = rel_w.reshape((b, self.num_heads, rel_w_dim1, rel_w_dim2, rel_w_dim3))?; + let attn_bias = rel_h.broadcast_add(&rel_w)?.reshape(( + b, + self.num_heads, + rel_h_dim1, + rel_h_dim2 * rel_w_dim3, + ))?; + let xs = eager_attention_forward( + &query_states, + &key_states, + &value_states, + None, + Some(&attn_bias), + self.scaling, + )?; + xs + } else { + eager_attention_forward( + &query_states, + &key_states, + &value_states, + None, + None, + self.scaling, + )? + }; + // (b, h*w, n_head, dim) + let xs = xs.reshape((b, h * w, ()))?.reshape((b, h, w, ()))?; + let xs = self.proj.forward(&xs)?; + Ok(xs) + } +} + +pub struct MLPBlock { + linear1: Linear, + linear2: Linear, + act: Activation, +} + +impl MLPBlock { + pub fn new( + vb: VarBuilder, + embedding_dim: usize, + mlp_dim: usize, + act: Activation, + ) -> Result { + let linear1 = linear(embedding_dim, mlp_dim, vb.pp("lin1"))?; + let linear2 = linear(mlp_dim, embedding_dim, vb.pp("lin2"))?; + Ok(Self { + linear1, + linear2, + act, + }) + } + pub fn forward(&self, xs: &Tensor) -> Result { + let xs = xs + .apply(&self.linear1)? + .apply(&self.act)? + .apply(&self.linear2)?; + Ok(xs) + } } pub struct Block { norm1: LayerNorm, attn: Attention, + norm2: LayerNorm, + mlp: MLPBlock, + window_size: usize, +} + +impl Block { + pub fn new( + vb: VarBuilder, + dim: usize, + num_heads: usize, + mlp_ratio: f32, + qkv_bias: bool, + eps: f64, + act: Activation, + use_rel_pos: bool, + rel_pos_zero_init: bool, + window_size: usize, + input_size: Option<(usize, usize)>, + ) -> Result { + let ln_config = LayerNormConfig { + eps, + remove_mean: true, // true for layernorm, false for RMSNorm + affine: true, // true for with bias, false for without bias + }; + let norm1 = layer_norm(dim, ln_config, vb.pp("norm1"))?; + let input_size = if window_size == 0 { + input_size + } else { + Some((window_size, window_size)) + }; + let attn = Attention::new( + vb.pp("attn"), + dim, + num_heads, + qkv_bias, + use_rel_pos, + input_size, + )?; + let norm2 = layer_norm(dim, ln_config, vb.pp("norm2"))?; + let mlp_dim = (dim as f32 * mlp_ratio) as usize; + let mlp = MLPBlock::new(vb.pp("mlp"), dim, mlp_dim, act)?; + Ok(Self { + norm1, + attn, + norm2, + mlp, + window_size, + }) + } + + pub fn window_partition( + &self, + x: &Tensor, + window_size: usize, + ) -> Result<(Tensor, (usize, usize))> { + let (b, h, w, c) = x.dims4()?; + let pad_h = (window_size - h % window_size) % window_size; + let pad_w = (window_size - w % window_size) % window_size; + let x = if pad_h > 0 || pad_w > 0 { + let x = x.pad_with_zeros(1, 0, pad_h)?; + let x = x.pad_with_zeros(2, 0, pad_w)?; + x + } else { + x.clone() + }; + let hp = h + pad_h; + let wp = w + pad_w; + let x = x.reshape(( + b, + hp / window_size, + window_size, + wp / window_size, + window_size, + c, + ))?; + let windows = x.permute((0, 1, 3, 2, 4, 5))?.contiguous()?.reshape(( + (), + window_size, + window_size, + c, + ))?; + Ok((windows, (hp, wp))) + } + + // pub fn window_unpartition( + // &self, + // x: &Tensor, + // window_size: usize, + // pad_hw: (usize, usize), + // hw: (usize, usize), + // ) -> Result { + + // } + + // pub fn forward(&self, xs: &Tensor) -> Result { + // let shortcut = xs.clone(); + // let xs = self.norm1.forward(xs)?; + // let xs = if self.window_size > 0 { + // let h = xs.dim(1)?; + // let w = xs.dim(2)?; + // let (x, (hp, wp)) = self.window_partition(&xs, self.window_size)?; + // let x = self.attn.forward(&x)?; + // } else { + // self.attn.forward(&xs)? + // }; + // } } pub struct ImageEncoderViT { diff --git a/src/utils/tensor_utils.rs b/src/utils/tensor_utils.rs index 2da0720..2850969 100644 --- a/src/utils/tensor_utils.rs +++ b/src/utils/tensor_utils.rs @@ -1,5 +1,6 @@ use anyhow::{Ok, Result, anyhow}; use candle_core::{D, DType, Device, IndexOp, Tensor, shape::Dim}; +use rocket::figment::value; pub fn prepare_causal_attention_mask( b_size: usize, @@ -8,10 +9,14 @@ pub fn prepare_causal_attention_mask( device: &Device, ) -> Result { // Sliding window mask? - let mask: Vec<_> = (0..tgt_len) - .flat_map(|i| (0..tgt_len).map(move |j| if i < j { f32::NEG_INFINITY } else { 0. })) - .collect(); - let mask = Tensor::from_slice(&mask, (tgt_len, tgt_len), device)?; + // let mask: Vec<_> = (0..tgt_len) + // .flat_map(|i| (0..tgt_len).map(move |j| if i < j { f32::NEG_INFINITY } else { 0. })) + // .collect(); + // let mask = Tensor::from_vec(mask, (tgt_len, tgt_len), device)?; + let arange = Tensor::arange(0u32, tgt_len as u32, device)?; + let arange = arange.unsqueeze(1)?.broadcast_as((tgt_len, tgt_len))?; + let upper_triangle = arange.t()?.lt(&arange)?.to_dtype(DType::F32)?; + let mask = upper_triangle.where_cond(&Tensor::new(f32::NEG_INFINITY, device)?, &Tensor::new(0f32, device)?)?; let mask = if seqlen_offset > 0 { let mask0 = Tensor::zeros((tgt_len, seqlen_offset), DType::F32, device)?; Tensor::cat(&[&mask0, &mask], D::Minus1)? @@ -346,3 +351,89 @@ pub fn mask_index_add(original: &Tensor, mask: &Tensor, add: &Tensor) -> Result< let xs = original.index_add(&visual_nonzero_index, add, 0)?; Ok(xs) } + +pub fn interpolate_linear( + t: &Tensor, + target_size: usize, + align_corner: Option, +) -> Result { + // t: [b, channels, features] + let shape = t.dims(); + let orig_size = shape[shape.len() - 1]; + if orig_size == target_size { + return Ok(t.clone()); + } + let mut reshaped = t.clone(); + if shape.len() != 3 { + let bs = shape[0]; + let channels = shape[1..shape.len() - 1].iter().product::(); + reshaped = reshaped.reshape((bs, channels, orig_size))?; + } + let (bs, channels, _) = reshaped.dims3()?; + let mut output = Tensor::zeros((bs, channels, target_size), t.dtype(), &t.device())?; + let coords = if orig_size == 1 { + vec![0f32; target_size] + } else { + let coords_vec = if let Some(align_) = align_corner + && align_ + { + (0..target_size) + .map(|i| i as f32 * (orig_size - 1) as f32 / (target_size - 1) as f32) + .collect() + } else { + (0..target_size) + .map(|i| { + let coord = (i as f32 + 0.5) * (orig_size as f32 / target_size as f32) - 0.5; + coord.max(0.0).min((orig_size-1) as f32) + }) + .collect() + }; + coords_vec + }; + + for b in 0..bs { + for c in 0..channels { + let input_slice = reshaped.i((b, c))?; + let mut out_i = Vec::new(); + for x_out in 0..target_size { + let coord = coords[x_out]; + let x0 = coord.floor() as usize; + let x1 = std::cmp::min(x0 + 1, orig_size - 1); + let weight = (coord - x0 as f32) as f64; + let value0 = input_slice.get(x0)?; + let value1 = input_slice.get(x1)?; + let interpolated = + (value0.affine(1.0 - weight, 0.0)? + value1.affine(weight, 0.0)?)?; + out_i.push(interpolated); + } + let out_i = Tensor::stack(&out_i, 0)?.unsqueeze(0)?.unsqueeze(0)?; + output = output.slice_assign(&[(b..b+1), (c..c+1), (0..target_size)], &out_i)?; + } + } + if shape.len() != 3 { + let mut new_shape = shape.to_vec(); + let last_dim = new_shape.len()-1; + new_shape[last_dim] = target_size; + output = output.reshape(new_shape)? + + } + output = output.contiguous()?; + Ok(output) +} + +pub fn index_select_2d(t: &Tensor, index: &Tensor) -> Result { + if t.rank() != 2 && index.rank() != 2 { + return Err(anyhow::anyhow!( + "t and index rank must be equal to 2" + )); + } + let mut res_vec = Vec::new(); + let index_dim0 = index.dim(0)?; + for i in 0..index_dim0 { + let index_i = index.i(i)?; + let rel_i = t.index_select(&index_i, 0)?; + res_vec.push(rel_i); + } + let res = Tensor::stack(&res_vec, 0)?; + Ok(res) +} \ No newline at end of file diff --git a/tests/messy_test.rs b/tests/messy_test.rs index 72d7062..1913b37 100644 --- a/tests/messy_test.rs +++ b/tests/messy_test.rs @@ -1,15 +1,38 @@ +use aha::utils::tensor_utils::{index_select_2d, interpolate_linear}; use anyhow::Result; use candle_core::{IndexOp, Tensor}; #[test] fn messy_test() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda messy_test -- --nocapture let device = &candle_core::Device::Cpu; - 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 t1 = Tensor::rand(0.0, 1.0, (1, 5, 5, 10), device)?; + let t2 = Tensor::rand(0.0, 1.0, (5, 8, 10), device)?; + let t2 = t2.t()?; + println!("t2: {:?}", t2); + let re = t1.broadcast_matmul(&t2)?; + println!("re: {:?}", re); + // let index = Tensor::arange(0u32, 10u32, device)?; + // let index_2d_vec = vec![index;5]; + // let index_2d = Tensor::stack(&index_2d_vec, 0)?; + // println!("index_2d: {}", index_2d); + // let t = Tensor::rand(0.0, 1.0, (20, 8), device)?; + // println!("t: {}", t); + // let res = index_select_2d(&t, &index_2d)?; + // println!("res: {}", res); + // let t = Tensor::arange(0.0, 10.0, device)? + // .unsqueeze(0)? + // .unsqueeze(0)?; + // println!("t: {}", t); + // let t_resized = interpolate_linear(&t, 20, None)?; + // println!("t_resized: {}", t_resized); + + // 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)?;