diff --git a/README.md b/README.md index 069390b..567ebfc 100644 --- a/README.md +++ b/README.md @@ -20,6 +20,7 @@ * Hunyuan-OCR - 腾讯混元光学文字识别模型 * PaddleOCR-VL - 百度飞桨光学文字识别模型 * VoxCPM1.5 - 面壁智能语音生成模型1.5版本 +* RMBG2.0 - RMBGv2.0由BRIA AI开发,供非商业用途使用。 ## 计划支持 我们持续扩展支持的模型列表,欢迎贡献! diff --git a/assets/img/gougou.jpg b/assets/img/gougou.jpg new file mode 100644 index 0000000..827d74a Binary files /dev/null and b/assets/img/gougou.jpg differ diff --git a/rmbg_0.png b/rmbg_0.png new file mode 100644 index 0000000..c76d6bf Binary files /dev/null and b/rmbg_0.png differ diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs index 315cb7f..9782441 100644 --- a/src/models/common/mod.rs +++ b/src/models/common/mod.rs @@ -1,8 +1,9 @@ use anyhow::Result; use candle_core::{D, Tensor}; use candle_nn::{ - Activation, Conv2d, Conv2dConfig, LayerNorm, LayerNormConfig, Linear, Module, RmsNorm, - VarBuilder, conv2d, conv2d_no_bias, layer_norm, linear, linear_no_bias, rms_norm, + Activation, BatchNorm, BatchNormConfig, Conv2d, Conv2dConfig, LayerNorm, LayerNormConfig, + Linear, Module, RmsNorm, VarBuilder, batch_norm, conv2d, conv2d_no_bias, layer_norm, linear, + linear_no_bias, rms_norm, }; use crate::{position_embed::rope::apply_rotary_pos_emb, utils::tensor_utils::repeat_kv}; @@ -511,3 +512,111 @@ pub fn get_layer_norm(vb: VarBuilder, eps: f64, dim: usize) -> Result let norm = layer_norm(dim, ln_config, vb)?; Ok(norm) } + +pub fn get_batch_norm(vb: VarBuilder, eps: f64, dim: usize) -> Result { + let bn_config = BatchNormConfig { + eps, + remove_mean: true, + affine: true, + momentum: 0.1, + }; + let norm = batch_norm(dim, bn_config, vb)?; + Ok(norm) +} + +pub fn deform_conv2d_kernel( + input: &Tensor, + weight: &Tensor, + bias: Option<&Tensor>, + offset: &Tensor, + mask: Option<&Tensor>, + stride: usize, + padding: usize, +) -> Result { + // 不考虑空洞卷积, bs = 1 + let (_, in_c, in_h, in_w) = input.dims4()?; + let (out_channel, _, ker_h, ker_w) = weight.dims4()?; + let out_h = ((in_h + 2 * padding - ker_h) / stride) + 1; + let out_w = ((in_w + 2 * padding - ker_w) / stride) + 1; + + let num_kernels = in_c * out_h * out_w; + let mask_vec = if let Some(mask) = mask { + Some(mask.squeeze(0)?.to_vec3::()?) + } else { + None + }; + let offset_vec = offset.squeeze(0)?.to_vec3::()?; + let input_vec = input.squeeze(0)?.to_vec3::()?; + let mut columns_vec = vec![vec![0.0f32; out_h * out_w]; in_c * ker_h * ker_w]; + for index in 0..num_kernels { + let out_x = index % out_w; + let out_y = (index / out_w) % out_h; + let in_c = index / (out_w * out_h); + let out_c = in_c * ker_h * ker_w; + + for i in 0..ker_h { + for j in 0..ker_w { + let mask_idx = i * ker_w + j; + let offset_idx = 2 * mask_idx; + let mask_value = if mask.is_some() { + mask_vec.as_ref().unwrap()[mask_idx][out_y][out_x] + } else { + 1.0 + }; + let offset_h = offset_vec[offset_idx][out_y][out_x]; + let offset_w = offset_vec[offset_idx + 1][out_y][out_x]; + let y = ((out_y * stride - padding) + i) as f32 + offset_h; + let x = ((out_x * stride - padding) + j) as f32 + offset_w; + let val = if y <= -1.0 || in_h as f32 <= y || x <= -1.0 || in_w as f32 <= x { + 0.0 + } else { + let h_low = y.floor(); + let w_low = x.floor(); + let h_high = h_low + 1.0; + let w_high = w_low + 1.0; + let lh = y - h_low; + let lw = x - w_low; + let hh = 1.0 - lh; + let hw = 1.0 - lw; + let w1 = hh * hw; + let w2 = hh * lw; + let w3 = lh * hw; + let w4 = lh * lw; + let v1 = if h_low >= 0.0 && w_low >= 0.0 { + input_vec[in_c][h_low as usize][w_low as usize] + } else { + 0.0 + }; + let v2 = if h_low >= 0.0 && w_high <= (in_w - 1) as f32 { + input_vec[in_c][h_low as usize][w_high as usize] + } else { + 0.0 + }; + let v3 = if h_high <= (in_h - 1) as f32 && w_low >= 0.0 { + input_vec[in_c][h_high as usize][w_low as usize] + } else { + 0.0 + }; + let v4 = if h_high <= (in_h - 1) as f32 && w_high <= (in_w - 1) as f32 { + input_vec[in_c][h_high as usize][w_high as usize] + } else { + 0.0 + }; + w1 * v1 + w2 * v2 + w3 * v3 + w4 * v4 + }; + columns_vec[out_c + i * ker_w + j][out_y * out_w + out_x] = mask_value * val; + } + } + } + + let columns = Tensor::new(columns_vec, weight.device())?; + let mut out = + weight + .flatten_from(1)? + .matmul(&columns)? + .reshape((1, out_channel, out_h, out_w))?; + if let Some(bias) = bias { + out = out.broadcast_add(bias)?; + } + Ok(out) +} diff --git a/src/models/mod.rs b/src/models/mod.rs index 9e54884..ed5a7f3 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -5,8 +5,8 @@ pub mod minicpm4; pub mod paddleocr_vl; pub mod qwen2_5vl; pub mod qwen3vl; -pub mod voxcpm; pub mod rmbg2_0; +pub mod voxcpm; use aha_openai_dive::v1::resources::chat::{ ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, diff --git a/src/models/rmbg2_0/config.rs b/src/models/rmbg2_0/config.rs deleted file mode 100644 index e69de29..0000000 diff --git a/src/models/rmbg2_0/generate.rs b/src/models/rmbg2_0/generate.rs index e69de29..1cff361 100644 --- a/src/models/rmbg2_0/generate.rs +++ b/src/models/rmbg2_0/generate.rs @@ -0,0 +1,80 @@ +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use anyhow::Result; +use candle_core::{DType, Device, Tensor}; +use candle_nn::VarBuilder; +use image::{Rgba, RgbaImage}; + +use crate::{ + models::rmbg2_0::model::BiRefNet, + utils::{ + find_type_files, get_device, get_dtype, + img_utils::{extract_images, float_tensor_to_dynamic_image, img_transform_with_resize}, + }, +}; + +pub struct RMBG2_0 { + model: BiRefNet, + h: u32, + w: u32, + img_mean: Tensor, + img_std: Tensor, + device: Device, + dtype: DType, +} + +impl RMBG2_0 { + pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { + let device = get_device(device); + let dtype = get_dtype(dtype, "float32"); + let model_list = find_type_files(path, "safetensors")?; + let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; + let model = BiRefNet::new(vb)?; + let img_mean = + Tensor::from_slice(&[0.485, 0.456, 0.406], (3, 1, 1), &device)?.to_dtype(dtype)?; + let img_std = + Tensor::from_slice(&[0.229, 0.224, 0.225], (3, 1, 1), &device)?.to_dtype(dtype)?; + Ok(Self { + model, + h: 1024, + w: 1024, + img_mean, + img_std, + device, + dtype, + }) + } + + pub fn generate(&self, mes: ChatCompletionParameters) -> Result> { + let imgs = extract_images(&mes)?; + let mut rmbg_png = vec![]; + for img in imgs { + let height = img.height(); + let width = img.width(); + let img_tensor = img_transform_with_resize( + &img, + self.h, + self.w, + &self.img_mean, + &self.img_std, + &self.device, + self.dtype, + )? + .unsqueeze(0)?; + let rmbg_img = self.model.forward(&img_tensor)?.squeeze(0)?; + let alpha_img = float_tensor_to_dynamic_image(&rmbg_img)?; + let alpha_img = + alpha_img.resize_exact(width, height, image::imageops::FilterType::CatmullRom); + let alpha_gray = alpha_img.to_luma8(); + let mut rgba_img = RgbaImage::new(width, height); + + // 遍历像素并组合 + for (x, y, pixel) in img.to_rgb8().enumerate_pixels() { + let alpha_value = alpha_gray.get_pixel(x, y).0[0]; + let rgba_pixel = Rgba([pixel.0[0], pixel.0[1], pixel.0[2], alpha_value]); + rgba_img.put_pixel(x, y, rgba_pixel); + } + rmbg_png.push(rgba_img); + } + Ok(rmbg_png) + } +} diff --git a/src/models/rmbg2_0/mod.rs b/src/models/rmbg2_0/mod.rs index 99bbd12..5671cb8 100644 --- a/src/models/rmbg2_0/mod.rs +++ b/src/models/rmbg2_0/mod.rs @@ -1 +1,2 @@ -pub mod model; \ No newline at end of file +pub mod generate; +pub mod model; diff --git a/src/models/rmbg2_0/model.rs b/src/models/rmbg2_0/model.rs index e648ec0..af74a18 100644 --- a/src/models/rmbg2_0/model.rs +++ b/src/models/rmbg2_0/model.rs @@ -1,8 +1,18 @@ -use anyhow::Result; -use candle_core::{D, Tensor}; -use candle_nn::{Activation, Conv2d, Dropout, LayerNorm, Module, VarBuilder}; +use anyhow::{Result, anyhow}; +use candle_core::{D, DType, Device, IndexOp, Shape, Tensor}; +use candle_nn::{ + Activation, BatchNorm, Conv2d, Init, LayerNorm, Linear, Module, ModuleT, VarBuilder, linear, + linear_no_bias, ops::sigmoid, +}; -use crate::models::common::{TwoLinearMLP, get_conv2d, get_layer_norm}; +use crate::{ + models::common::{ + TwoLinearMLP, deform_conv2d_kernel, get_batch_norm, get_conv2d, get_layer_norm, + }, + utils::tensor_utils::{ + get_equal_mask, index_select_2d, interpolate_bilinear, split_tensor_with_size, + }, +}; struct PatchEmbed { proj: Conv2d, @@ -16,19 +26,16 @@ impl PatchEmbed { vb: VarBuilder, in_chans: usize, embed_dim: usize, - kernel_size: usize, - stride: usize, - padding: usize, + patch_size: usize, patch_norm: bool, ) -> Result { - let patch_size = kernel_size; let proj = get_conv2d( vb.pp("proj"), in_chans, embed_dim, - kernel_size, - padding, - stride, + patch_size, + 0, + patch_size, 1, 1, true, @@ -47,7 +54,7 @@ impl PatchEmbed { } pub fn forward(&self, xs: &Tensor) -> Result { - let (bs, _, h, w) = xs.dims4()?; + let (_, _, h, w) = xs.dims4()?; let mut xs = xs.clone(); if w % self.patch_size != 0 { xs = xs.pad_with_zeros(3, 0, self.patch_size - w % self.patch_size)?; @@ -58,9 +65,9 @@ impl PatchEmbed { xs = self.proj.forward(&xs)?; if self.norm.is_some() { let (_, _, ph, pw) = xs.dims4()?; - xs = xs.flatten(2, 3)?.transpose(1, 2)?; + xs = xs.flatten_from(2)?.transpose(1, 2)?; xs = self.norm.as_ref().unwrap().forward(&xs)?; - xs = xs.transpose(1, 2)?.reshape((bs, self.embed_dim, ph, pw))?; + xs = xs.transpose(1, 2)?.reshape(((), self.embed_dim, ph, pw))?; } Ok(xs) } @@ -68,13 +75,10 @@ impl PatchEmbed { pub struct WindowAttention { num_heads: usize, - // head_dim: usize, + relative_position_bias: Tensor, qkv: Linear, proj: Linear, scaling: f64, - use_rel_pos: bool, - rel_pos_h: Option, - rel_pos_w: Option, } impl WindowAttention { @@ -83,8 +87,7 @@ impl WindowAttention { dim: usize, num_heads: usize, qkv_bias: bool, - window_size: usize, /* use_rel_pos: bool, - * input_size: Option<(usize, usize)>, */ + window_size: (usize, usize), ) -> Result { let head_dim = dim / num_heads; let scaling = 1.0 / (head_dim as f64).sqrt(); @@ -94,154 +97,133 @@ impl WindowAttention { linear_no_bias(dim, dim * 3, vb.pp("qkv"))? }; let proj = linear(dim, dim, vb.pp("proj"))?; - let mut rel_pos_h = None; - let mut rel_pos_w = None; - if use_rel_pos { - if input_size.is_none() { - return Err(anyhow::anyhow!( - "Input size must be provided if using relative positional encoding." - )); - } - let input_size = input_size.unwrap(); - let h_len = 2 * input_size.0 - 1; - let w_len = 2 * input_size.1 - 1; - rel_pos_h = Some(vb.get_with_hints((h_len, head_dim), "rel_pos_h", Init::Const(0.))?); - rel_pos_w = Some(vb.get_with_hints((w_len, head_dim), "rel_pos_w", Init::Const(0.))?); - } + let relative_position_bias_table = vb.get_with_hints( + ((2 * window_size.0 - 1) * (2 * window_size.1 - 1), num_heads), + "relative_position_bias_table", + Init::Const(0.), + )?; //2*Wh-1 * 2*Ww-1, nH + let coords_h = Tensor::arange(0f32, window_size.0 as f32, vb.device())? + .unsqueeze(1)? + .broadcast_as(window_size)?; + let coords_w = Tensor::arange(0f32, window_size.1 as f32, vb.device())? + .unsqueeze(0)? + .broadcast_as(window_size)?; + + let coords = Tensor::stack(&[coords_h, coords_w], 0)?.flatten_from(1)?; // (2, wh, ww) + let coords1 = coords.unsqueeze(2)?; + let coords2 = coords.unsqueeze(1)?; + let relative_coords = coords1 + .broadcast_sub(&coords2)? + .permute((1, 2, 0))? + .contiguous()?; // (wh*ww, wh*ww, 2) + let relative_coords_0 = relative_coords + .i((.., .., 0))? + .affine(1.0, window_size.0 as f64 - 1.0)?; + let relative_coords_1 = relative_coords + .i((.., .., 1))? + .affine(1.0, window_size.1 as f64 - 1.0)?; + let relative_coords_0 = relative_coords_0.affine(2.0 * window_size.1 as f64 - 1.0, 0.0)?; + let relative_position_index = relative_coords_0 + .add(&relative_coords_1)? + .to_dtype(candle_core::DType::U32)?; + let relative_position_bias = + index_select_2d(&relative_position_bias_table, &relative_position_index)?; Ok(Self { num_heads, - // head_dim, + relative_position_bias, qkv, proj, scaling, - use_rel_pos, - rel_pos_h, - rel_pos_w, }) } - 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_t = rel_pos - .to_dtype(candle_core::DType::F32)? - .t()? - .unsqueeze(0)? - .contiguous()?; - let rel_pos_resized = interpolate_linear_1d(&rel_pos_t, max_rel_dist, None)?; - rel_pos_resized - .squeeze(0)? - .t()? - .contiguous()? - .to_dtype(rel_pos.dtype())? - } 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) - } - - 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)?; // (q_h, k_h, dim) - let rw = self.get_rel_pos(q_w, k_w, rel_pos_w)?; // (q_w, k_w, dim) - let (b, _, dim) = q.dims3()?; - let r_q = q.reshape((b, q_h, q_w, dim))?.contiguous()?; - let r_q_ = r_q.unsqueeze(D::Minus2)?; // (b, q_h, q_w, 1, dim) - // rel_h = torch.einsum("bhwc,hkc->bhwk", r_q, Rh) - // rel_w = torch.einsum("bhwc,wkc->bhwk", r_q, Rw) - let rh_ = rh.unsqueeze(1)?.unsqueeze(0)?; // (1, h, 1, k, dim) - let rel_h = r_q_.broadcast_mul(&rh_)?.sum(D::Minus1)?; - let rw_ = rw.unsqueeze(0)?.unsqueeze(0)?; // (1, 1, w, k, dim) - let rel_w = r_q_.broadcast_mul(&rw_)?.sum(D::Minus1)?; - 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(&self, xs: &Tensor, attn_mask: Option<&Tensor>) -> Result { - let (b, h, w, _) = xs.dims4()?; + let (b, seq_len, _) = xs.dims3()?; // (3, B, n_head, h*w, head_dim) let qkv = self .qkv .forward(xs)? - .reshape((b, h * w, 3, self.num_heads, ()))? + .reshape((b, seq_len, 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, - ))?; - eager_attention_forward( - &query_states, - &key_states, - &value_states, - None, - Some(&attn_bias), - self.scaling, - )? - } else { - eager_attention_forward( - &query_states, - &key_states, - &value_states, - None, - None, - self.scaling, - )? + let attn_bias = self + .relative_position_bias + .permute((2, 0, 1))? + .contiguous()? + .unsqueeze(0)?; + let query_states = (query_states * self.scaling)?; + + let attn_weights = query_states.matmul(&key_states.transpose(D::Minus2, D::Minus1)?)?; + let attn_weights = attn_weights.broadcast_add(&attn_bias)?; + let attn_weights = match attn_mask { + None => attn_weights, + Some(mask) => { + let nw: usize = mask.dim(0)?; + let attn_weights = attn_weights + .reshape((b / nw, nw, self.num_heads, seq_len, seq_len))? + .broadcast_add( + &mask + .unsqueeze(1)? + .unsqueeze(0)? + .to_dtype(attn_weights.dtype())?, + )?; + + attn_weights.reshape(((), self.num_heads, seq_len, seq_len))? + } }; + let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?; + let attn_output = attn_weights.matmul(&value_states)?; + + //(b, n_head, seq_len, dim) -> (b, seq_len, n_head, dim) + let xs = attn_output.transpose(1, 2)?.contiguous()?; // (b, h*w, n_head, dim) - let xs = xs.reshape((b, h * w, ()))?.reshape((b, h, w, ()))?; + let xs = xs.reshape((b, seq_len, ()))?; let xs = self.proj.forward(&xs)?; Ok(xs) } } +fn window_partition(x: &Tensor, window_size: usize) -> Result { + let (b, h, w, c) = x.dims4()?; + + let x = x.reshape(( + b, + h / window_size, + window_size, + w / window_size, + window_size, + c, + ))?; + let windows = + x.permute((0, 1, 3, 2, 4, 5))? + .contiguous()? + .reshape(((), window_size, window_size, c))?; + Ok(windows) +} + +fn window_reverse(windows: &Tensor, window_size: usize, pad_hw: (usize, usize)) -> Result { + let (hp, wp) = pad_hw; + let b = windows.dim(0)? / (hp * wp / window_size / window_size); + let last_dim = windows.dim(D::Minus1)?; + let x = windows.reshape(&[ + b, + hp / window_size, + wp / window_size, + window_size, + window_size, + last_dim, + ])?; + let x = x + .permute((0, 1, 3, 2, 4, 5))? + .contiguous()? + .reshape((b, hp, wp, ()))?; + Ok(x) +} + struct SwinTransformerBlock { norm1: LayerNorm, attn: WindowAttention, @@ -264,7 +246,13 @@ impl SwinTransformerBlock { ) -> Result { let norm1 = get_layer_norm(vb.pp("norm1"), 1e-5, dim)?; - let attn = WindowAttention::new(vb.pp("attn"), dim, num_heads, qkv_bias, window_size)?; + let attn = WindowAttention::new( + vb.pp("attn"), + dim, + num_heads, + qkv_bias, + (window_size, window_size), + )?; let norm2 = get_layer_norm(vb.pp("norm2"), 1e-5, dim)?; let mlp_dim = (dim as f32 * mlp_ratio) as usize; let mlp = TwoLinearMLP::new(vb.pp("mlp"), dim, mlp_dim, act, true, "fc1", "fc2")?; @@ -278,50 +266,6 @@ impl SwinTransformerBlock { }) } - pub fn window_partition(&self, x: &Tensor, window_size: usize) -> Result { - let (b, h, w, c) = x.dims4()?; - - let x = x.reshape(( - b, - h / window_size, - window_size, - w / window_size, - window_size, - c, - ))?; - let windows = x.permute((0, 1, 3, 2, 4, 5))?.contiguous()?.reshape(( - (), - window_size, - window_size, - c, - ))?; - Ok(windows) - } - - pub fn window_unpartition( - &self, - windows: &Tensor, - window_size: usize, - pad_hw: (usize, usize), - ) -> Result { - let (hp, wp) = pad_hw; - let b = windows.dim(0)? / (hp * wp / window_size / window_size); - let last_dim = windows.dim(D::Minus1)?; - let x = windows.reshape(&[ - b, - hp / window_size, - wp / window_size, - window_size, - window_size, - last_dim, - ])?; - let x = x - .permute((0, 1, 3, 2, 4, 5))? - .contiguous()? - .reshape((b, hp, wp, ()))?; - Ok(x) - } - pub fn forward( &self, xs: &Tensor, @@ -330,6 +274,11 @@ impl SwinTransformerBlock { w: usize, ) -> Result { let (b, seq_len, c) = xs.dims3()?; + assert_eq!( + seq_len, + h * w, + "swin transformer block sq_len not equal to h*w" + ); let shortcut = xs.clone(); let xs = self.norm1.forward(xs)?; let xs = xs.reshape((b, h, w, c))?; @@ -348,10 +297,10 @@ impl SwinTransformerBlock { } else { (xs, None) }; - let xs = self.window_partition(&shifted_x, self.window_size)?; + let xs = window_partition(&shifted_x, self.window_size)?; let xs = xs.reshape(((), self.window_size * self.window_size, c))?; let xs = self.attn.forward(&xs, attn_mask)?; - let xs = self.window_unpartition(&xs, self.window_size, (hp, wp))?; + let xs = window_reverse(&xs, self.window_size, (hp, wp))?; let mut xs = if self.shift_size > 0 { xs.roll(self.shift_size as i32, 1)? .roll(self.shift_size as i32, 2)? @@ -359,25 +308,1097 @@ impl SwinTransformerBlock { xs }; if pad_h > 0 || pad_w > 0 { - xs = xs.i((.., 0..h, 0..w, ..))? + xs = xs.i((.., 0..h, 0..w, ..))?.contiguous()?; } + let xs = xs.reshape((b, h * w, c))?; let x = shortcut.add(&xs)?; let x = x.add(&self.mlp.forward(&self.norm2.forward(&x)?)?)?; Ok(x) } } -struct PatchMerging {} +struct PatchMerging { + reduction: Linear, + norm: LayerNorm, +} + +impl PatchMerging { + pub fn new(vb: VarBuilder, dim: usize) -> Result { + let reduction = linear_no_bias(4 * dim, 2 * dim, vb.pp("reduction"))?; + let norm = get_layer_norm(vb.pp("norm"), 1e-5, 4 * dim)?; + Ok(Self { reduction, norm }) + } + + pub fn forward(&self, xs: &Tensor, h: usize, w: usize) -> Result { + let (b, l, c) = xs.dims3()?; + assert_eq!(l, h * w, "input feature has wrong size"); + let mut xs = xs.reshape((b, h, w, c))?; + let pad_input = (h % 2 == 1) || (w % 2 == 1); + if pad_input { + xs = xs + .pad_with_zeros(2, 0, w % 2)? + .pad_with_zeros(1, 0, h % 2)?; + } + let shape = Shape::from_dims(&[b, h / 2, 2, w / 2, 2, c]); + let xs = xs.reshape(shape)?; + let x0 = xs.i((.., .., 0, .., 0, ..))?; + let x1 = xs.i((.., .., 1, .., 0, ..))?; + let x2 = xs.i((.., .., 0, .., 1, ..))?; + let x3 = xs.i((.., .., 1, .., 1, ..))?; + let xs = Tensor::cat(&[x0, x1, x2, x3], D::Minus1)?; + let xs = xs.reshape((b, (), 4 * c))?; + let xs = self.norm.forward(&xs)?; + let xs = self.reduction.forward(&xs)?; + Ok(xs) + } +} struct BasicLayer { - windows_size: usize, + window_size: usize, shift_size: usize, blocks: Vec, downsample: Option, } + +impl BasicLayer { + pub fn new( + vb: VarBuilder, + dim: usize, + depth: usize, + num_heads: usize, + window_size: usize, + mlp_ratio: f32, + qkv_bias: bool, + downsample: bool, + ) -> Result { + let shift_size = window_size / 2; + let mut blocks = vec![]; + let vb_blocks = vb.pp("blocks"); + for i in 0..depth { + let block_shift_size = if i % 2 == 0 { 0usize } else { shift_size }; + + let block = SwinTransformerBlock::new( + vb_blocks.pp(i), + dim, + num_heads, + mlp_ratio, + qkv_bias, + Activation::Gelu, + window_size, + block_shift_size, + )?; + blocks.push(block); + } + let downsample = if downsample { + Some(PatchMerging::new(vb.pp("downsample"), dim)?) + } else { + None + }; + Ok(Self { + window_size, + shift_size, + blocks, + downsample, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + h: usize, + w: usize, + ) -> Result<(Tensor, usize, usize, Tensor, usize, usize)> { + let hp = (h as f32 / self.window_size as f32).ceil() as usize * self.window_size; + let wp = (w as f32 / self.window_size as f32).ceil() as usize * self.window_size; + let mut img_mask = Tensor::zeros((1, hp, wp, 1), xs.dtype(), xs.device())?; + let h_slices = [ + (0usize, hp - self.window_size), + (hp - self.window_size, hp - self.shift_size), + (hp - self.shift_size, hp), + ]; + let w_slices = [ + (0usize, wp - self.window_size), + (wp - self.window_size, wp - self.shift_size), + (wp - self.shift_size, wp), + ]; + let mut cnt = 0f64; + for (h_start, h_end) in h_slices { + for (w_start, w_end) in w_slices { + let mask_value = Tensor::zeros( + (1, h_end - h_start, w_end - w_start, 1), + xs.dtype(), + xs.device(), + )? + .affine(1.0, cnt)?; + img_mask = img_mask.slice_assign( + &[(0..1), (h_start..h_end), (w_start..w_end), (0..1)], + &mask_value, + )?; + cnt += 1.0; + } + } + + let mask_windows = window_partition(&img_mask, self.window_size)?; + let mask_windows = mask_windows.reshape(((), self.window_size * self.window_size))?; + let attn_mask = mask_windows + .unsqueeze(1)? + .broadcast_sub(&mask_windows.unsqueeze(2)?)?; + let equal_zero_mask = get_equal_mask(&attn_mask, 0)?; + let attn_mask = equal_zero_mask.where_cond( + &Tensor::new(0f32, xs.device())?.broadcast_as(equal_zero_mask.shape())?, + &Tensor::new(-100f32, xs.device())?.broadcast_as(equal_zero_mask.shape())?, + )?; + let mut xs = xs.clone(); + for block in &self.blocks { + xs = block.forward(&xs, Some(&attn_mask), h, w)?; + } + let (xs_down, wh, ww) = match self.downsample.as_ref() { + Some(down) => { + let xs_down = down.forward(&xs, h, w)?; + // let wh = (h + 1) / 2; + // let ww = (w + 1) / 2; + let wh = h.div_ceil(2); + let ww = w.div_ceil(2); + (xs_down, wh, ww) + } + None => (xs.clone(), h, w), + }; + Ok((xs, h, w, xs_down, wh, ww)) + } +} + pub struct SwinTransformer { patch_embed: PatchEmbed, - pos_drop: Dropout, + num_layers: usize, + // pos_drop: Dropout, layers: Vec, norms: Vec, + out_indices: Vec, + num_features: Vec, +} + +impl SwinTransformer { + pub fn new( + vb: VarBuilder, + patch_size: usize, + in_channels: usize, + embed_dim: usize, + depths: Vec, + num_heads: Vec, + window_size: usize, + mlp_ratio: f32, + qkv_bias: bool, + patch_norm: bool, + out_indices: Vec, + ) -> Result { + let patch_embed = PatchEmbed::new( + vb.pp("patch_embed"), + in_channels, + embed_dim, + patch_size, + patch_norm, + )?; + let num_layers = depths.len(); + let mut layers = vec![]; + let vb_layers = vb.pp("layers"); + let mut num_features = vec![]; + for i in 0..num_layers { + let downsample = i < num_layers - 1; + let dim_i = embed_dim * 2usize.pow(i as u32); + num_features.push(dim_i); + let layer_i = BasicLayer::new( + vb_layers.pp(i), + dim_i, + depths[i], + num_heads[i], + window_size, + mlp_ratio, + qkv_bias, + downsample, + )?; + layers.push(layer_i); + } + let mut norms = vec![]; + for i in out_indices.clone() { + let layer_i = get_layer_norm(vb.pp(format!("norm{i}")), 1e-5, num_features[i])?; + norms.push(layer_i); + } + Ok(Self { + num_layers, + patch_embed, + layers, + norms, + out_indices, + num_features, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result> { + let xs = self.patch_embed.forward(xs)?; + let (_, _, mut wh, mut ww) = xs.dims4()?; + let mut outs = vec![]; + let mut xs = xs.flatten_from(2)?.transpose(1, 2)?; + let mut norm_idx = 0; + for i in 0..self.num_layers { + let layer = &self.layers[i]; + let (x_out, h, w, xs_, wh_, ww_) = layer.forward(&xs, wh, ww)?; + xs = xs_.clone(); + wh = wh_; + ww = ww_; + if self.out_indices.contains(&i) { + let norm_layer = &self.norms[norm_idx]; + norm_idx += 1; + let x_out = norm_layer.forward(&x_out)?; + let out = x_out + .reshape(((), h, w, self.num_features[i]))? + .permute((0, 3, 1, 2))? + .contiguous()?; + outs.push(out); + } + } + Ok(outs) + } +} + +#[allow(unused)] +struct DeformableConv2d { + offset_conv: Conv2d, + modulator_conv: Conv2d, + regular_conv: Conv2d, + stride: usize, + padding: usize, + ks: usize, +} + +#[allow(unused)] +impl DeformableConv2d { + pub fn new( + vb: VarBuilder, + in_c: usize, + out_c: usize, + kernel_size: usize, + stride: usize, + padding: usize, + bias: bool, + ) -> Result { + let offset_conv = get_conv2d( + vb.pp("offset_conv"), + in_c, + 2 * kernel_size * kernel_size, + kernel_size, + padding, + stride, + 1, + 1, + true, + )?; + + let modulator_conv = get_conv2d( + vb.pp("modulator_conv"), + in_c, + kernel_size * kernel_size, + kernel_size, + padding, + stride, + 1, + 1, + true, + )?; + + let regular_conv = get_conv2d( + vb.pp("regular_conv"), + in_c, + out_c, + kernel_size, + 0, + kernel_size, + 1, + 1, + bias, + )?; + Ok(Self { + offset_conv, + modulator_conv, + regular_conv, + stride, + padding, + ks: kernel_size, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + self.forward_use_kernel(xs) + // if self.ks > 1 { + // self.forward_use_kernel(xs) + // } else { + // self.forward_use_tensor(xs) + // } + } + pub fn forward_use_kernel(&self, xs: &Tensor) -> Result { + let offset = self.offset_conv.forward(xs)?; // (b, 2*k*k, out_h, out_w) + + let modulator = sigmoid(&self.modulator_conv.forward(xs)?)? + .affine(2.0, 0.0)? + .contiguous()?; + let out = deform_conv2d_kernel( + xs, + self.regular_conv.weight(), + self.regular_conv.bias(), + &offset, + Some(&modulator), + self.stride, + self.padding, + )?; + Ok(out) + } + + pub fn forward_use_tensor(&self, xs: &Tensor) -> Result { + let offset = self.offset_conv.forward(xs)?; // (b, 2*k*k, out_h, out_w) + + let modulator = sigmoid(&self.modulator_conv.forward(xs)?)? + .affine(2.0, 0.0)? + .contiguous()?; + let n = offset.dim(1)? / 2; + + let xs = if self.padding > 0 { + xs.pad_with_zeros(2, self.padding, self.padding)? + .pad_with_zeros(3, self.padding, self.padding)? + } else { + xs.clone() + }; + let offset = if self.ks > 3 { + offset.to_device(&Device::Cpu)? + } else { + offset + }; + let p = self.get_p(&offset)?; + // drop(offset); + // (b, h, w, 2n) + let p = p.permute((0, 2, 3, 1))?.contiguous()?; + let q_lt = p.floor()?; + let q_rb = (&q_lt + 1.0)?; + let (_, _, in_h, in_w) = xs.dims4()?; + let in_h = in_h as f64; + let in_w = in_w as f64; + + // 分开处理x和y坐标 + let p_x = p.narrow(3, 0, n)?.clamp(0.0, in_h - 1.0)?; + let p_y = p.narrow(3, n, n)?.clamp(0.0, in_w - 1.0)?; + // drop(p); + let q_lt_x = q_lt.narrow(3, 0, n)?.clamp(0.0, in_h - 1.0)?; + let q_lt_y = q_lt.narrow(3, n, n)?.clamp(0.0, in_w - 1.0)?; + let q_rb_x = q_rb.narrow(3, 0, n)?.clamp(0.0, in_h - 1.0)?; + let q_rb_y = q_rb.narrow(3, n, n)?.clamp(0.0, in_w - 1.0)?; + // drop(q_lt); + // drop(q_rb); + // 转换为整数索引 + let q_lt_x_idx = q_lt_x.to_dtype(DType::U32)?; + let q_lt_y_idx = q_lt_y.to_dtype(DType::U32)?; + let q_rb_x_idx = q_rb_x.to_dtype(DType::U32)?; + let q_rb_y_idx = q_rb_y.to_dtype(DType::U32)?; + + // 计算双线性权重 + let p_sub_lt_x = (&p_x - &q_lt_x)?; + let one_sub_lt_x = (1.0 - &p_sub_lt_x)?; + let p_sub_lt_y = (&p_y - &q_lt_y)?; + let one_sub_lt_y = (1.0 - &p_sub_lt_y)?; + // drop(q_lt_x); + // drop(q_lt_y); + // drop(q_rb_x); + // drop(q_rb_y); + let g_lt = (&one_sub_lt_x * &one_sub_lt_y)?; + let g_rb = (&p_sub_lt_x * &p_sub_lt_y)?; + let g_lb = (&one_sub_lt_x * &p_sub_lt_y)?; + let g_rt = (&p_sub_lt_x * &one_sub_lt_y)?; + // drop(p_sub_lt_x); + // drop(one_sub_lt_x); + // drop(p_sub_lt_y); + // drop(one_sub_lt_y); + + let xs = if self.ks > 3 { + xs.to_device(&Device::Cpu)? + } else { + xs + }; + // 采样四个角点的特征 + let x_q_lt = self.get_x_q(&xs, &q_lt_x_idx, &q_lt_y_idx)?; + let x_q_rb = self.get_x_q(&xs, &q_rb_x_idx, &q_rb_y_idx)?; + let x_q_lb = self.get_x_q(&xs, &q_lt_x_idx, &q_rb_y_idx)?; + let x_q_rt = self.get_x_q(&xs, &q_rb_x_idx, &q_lt_y_idx)?; + // drop(q_lt_x_idx); + // drop(q_lt_y_idx); + // drop(q_rb_x_idx); + // drop(q_rb_y_idx); + // 双线性插值 + let x_offset = g_lt.unsqueeze(1)?.broadcast_mul(&x_q_lt)?; + // drop(g_lt); + // drop(x_q_lt); + let x_offset = x_offset.add(&g_rb.unsqueeze(1)?.broadcast_mul(&x_q_rb)?)?; + // drop(g_rb); + // drop(x_q_rb); + let x_offset = x_offset.add(&g_lb.unsqueeze(1)?.broadcast_mul(&x_q_lb)?)?; + // drop(g_lb); + // drop(x_q_lb); + let x_offset = x_offset.add(&g_rt.unsqueeze(1)?.broadcast_mul(&x_q_rt)?)?; + // drop(g_rt); + // drop(x_q_rt); + // (bs, n, h, w) -> (bs, h, w, n) -> (bs, 1, h, w, n) + let m = modulator.permute((0, 2, 3, 1))?.unsqueeze(1)?; + let x_offset = x_offset.to_device(m.device())?.broadcast_mul(&m)?; + let x_offset = self.reshape_x_offset(&x_offset, self.ks)?; + let xs = self.regular_conv.forward(&x_offset)?; + Ok(xs) + } + fn reshape_x_offset(&self, xs: &Tensor, ks: usize) -> Result { + let (b, c, h, w, _) = xs.dims5()?; + + let xs = xs.reshape((b, c, h, w, ks, ks))?; + let xs = xs.permute((0, 1, 2, 4, 3, 5))?; + let xs = xs.reshape((b, c, h * ks, w * ks))?; + let x_offset = xs.contiguous()?; + + Ok(x_offset) + } + + fn get_x_q(&self, xs: &Tensor, q_x: &Tensor, q_y: &Tensor) -> Result { + let (b, h, w, n) = q_x.dims4()?; + let padded_w = xs.dim(3)?; + let c = xs.dim(1)?; + + // 展平输入: (b, c, H, W) -> (b, c, H*W) + let xs_flat = xs.flatten_from(2)?; + + // 计算索引 + let index = q_x.affine(padded_w as f64, 0.0)?.add(q_y)?; // (b, h, w, n) + + // 扩展维度以匹配通道数 + let index = index + .unsqueeze(1)? + .expand((b, c, h, w, n))? + .flatten_from(2)?; // (b, c, h*w*n) + + // 收集特征 + let xs = xs_flat.gather(&index, 2)?.reshape((b, c, h, w, n))?; + Ok(xs) + } + + fn get_p_n(&self, n: usize, dtype: DType, device: &Device) -> Result { + let ks = self.ks as f32; + let range = Tensor::arange_step(-(ks - 1.0) / 2.0, (ks - 1.0) / 2.0 + 1.0, 1.0, device)?; + // 假设 ks=3 + // [(-1, -1), (-1, 0), (-1, 1) + // (0, -1), (0, 0), (0, 1) + // (1, -1), (1, 0), (1, 1)] + // range: [-1, 0, 1] + // unsqueeze(1) -> [[-1], [0], [1]] + // broadcase_as(3, 3) -> [[-1, -1, -1], [0, 0, 0], [1, 1, 1]] + // flatten_all -> [-1, -1, -1, 0, 0, 0, 1, 1, 1] + let p_n_x = range + .unsqueeze(1)? + .broadcast_as((self.ks, self.ks))? + .flatten_all()?; + // range: [-1, 0, 1] + // unsqueeze(0) -> [[-1, 0, 1]] + // broadcase_as(3, 3) -> [[-1, 0, 1], [-1, 0, 1], [-1, 0, 1]] + // flatten_all -> [-1, 0, 1, -1, 0, 1, -1, 0, 1] + let p_n_y = range + .unsqueeze(0)? + .broadcast_as((self.ks, self.ks))? + .flatten_all()?; + let p = Tensor::cat(&[p_n_x, p_n_y], 0)? + .reshape((1, 2 * n, 1, 1))? + .to_dtype(dtype)? + .contiguous()?; + Ok(p) + } + + fn get_p_0( + &self, + h: usize, + w: usize, + n: usize, + dtype: DType, + device: &Device, + ) -> Result { + // 假设 in featuremap h=w=5, padding=1, hp=wp=7 + // out featuremap h=w=5, + // padding 后的 in featuremap + // [(0, 0), (0, 1), (0, 2), (0, 3), (0, 4), (0, 5), (0, 6) + // (1, 0), (1, 1), (1, 2), (1, 3), (1, 4), (1, 5), (1, 6) + // (2, 0), (2, 1), (2, 2), (2, 3), (2, 4), (2, 5), (2, 6) + // (3, 0), (3, 1), (3, 2), (3, 3), (3, 4), (3, 5), (3, 6) + // (4, 0), (4, 1), (4, 2), (4, 3), (4, 4), (4, 5), (4, 6) + // (5, 0), (5, 1), (5, 2), (5, 3), (5, 4), (5, 5), (5, 6) + // (6, 0), (6, 1), (6, 2), (6, 3), (6, 4), (6, 5), (6, 6) + let start = self.padding as f32; + // let start = 0.0f32; + let p_0_x = Tensor::arange_step( + start, + start + h as f32 * self.stride as f32, + self.stride as f32, + device, + )? + .unsqueeze(1)? + .broadcast_as((h, w))? + .reshape((1, 1, h, w))? + .repeat((1, n, 1, 1))?; + let p_0_y = Tensor::arange_step( + start, + start + w as f32 * self.stride as f32, + self.stride as f32, + device, + )? + .unsqueeze(0)? + .broadcast_as((h, w))? + .reshape((1, 1, h, w))? + .repeat((1, n, 1, 1))?; + let p_0 = Tensor::cat(&[p_0_x, p_0_y], 1)? + .to_dtype(dtype)? + .contiguous()?; + Ok(p_0) + } + fn get_p(&self, offset: &Tensor) -> Result { + let (_, n, h, w) = offset.dims4()?; + let n = n / 2; + // (1, 2n, 1, 1) + let p_n = self.get_p_n(n, offset.dtype(), offset.device())?; + // (1, 2n, h, w) + let p_0 = self.get_p_0(h, w, n, offset.dtype(), offset.device())?; + let p = p_0 + .broadcast_add(&p_n)? + .broadcast_add(offset)? + .contiguous()?; + Ok(p) + } +} + +struct _ASPPModuleDeformable { + atrous_conv: DeformableConv2d, + bn: BatchNorm, +} + +impl _ASPPModuleDeformable { + pub fn new( + vb: VarBuilder, + in_c: usize, + out_c: usize, + kernel_size: usize, + padding: usize, + ) -> Result { + let atrous_conv = DeformableConv2d::new( + vb.pp("atrous_conv"), + in_c, + out_c, + kernel_size, + 1, + padding, + false, + )?; + let bn = get_batch_norm(vb.pp("bn"), 1e-5, out_c)?; + Ok(Self { atrous_conv, bn }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let xs = self.atrous_conv.forward(xs)?; + let xs = self.bn.forward_t(&xs, false)?.relu()?; + Ok(xs) + } +} + +struct ASPPDeformable { + aspp1: _ASPPModuleDeformable, + // aspp_deforms: Vec<_ASPPModuleDeformable>, + aspp_deforms_0: _ASPPModuleDeformable, + aspp_deforms_1: _ASPPModuleDeformable, + aspp_deforms_2: _ASPPModuleDeformable, + // avgpool2d + conv2d + BatchNorm2d + relu + global_avg_pool_1: Conv2d, + global_avg_pool_2: BatchNorm, + conv1: Conv2d, + bn1: BatchNorm, +} + +impl ASPPDeformable { + pub fn new( + vb: VarBuilder, + in_c: usize, + out_c: usize, + parallel_block_sizes: Vec, + ) -> Result { + let in_channelster = 256; + let aspp1 = _ASPPModuleDeformable::new(vb.pp("aspp1"), in_c, in_channelster, 1, 0)?; + let vb_aspp_deforms = vb.pp("aspp_deforms"); + let aspp_deforms_0 = _ASPPModuleDeformable::new( + vb_aspp_deforms.pp(0), + in_c, + in_channelster, + parallel_block_sizes[0], + parallel_block_sizes[0] / 2, + )?; + let aspp_deforms_1 = _ASPPModuleDeformable::new( + vb_aspp_deforms.pp(1), + in_c, + in_channelster, + parallel_block_sizes[1], + parallel_block_sizes[1] / 2, + )?; + let aspp_deforms_2 = _ASPPModuleDeformable::new( + vb_aspp_deforms.pp(2), + in_c, + in_channelster, + parallel_block_sizes[2], + parallel_block_sizes[2] / 2, + )?; + let global_avg_pool_1 = get_conv2d( + vb.pp("global_avg_pool.1"), + in_c, + in_channelster, + 1, + 0, + 1, + 1, + 1, + false, + )?; + let global_avg_pool_2 = get_batch_norm(vb.pp("global_avg_pool.2"), 1e-5, in_channelster)?; + let conv1 = get_conv2d( + vb.pp("conv1"), + in_channelster * (2 + parallel_block_sizes.len()), + out_c, + 1, + 0, + 1, + 1, + 1, + false, + )?; + let bn1 = get_batch_norm(vb.pp("bn1"), 1e-5, out_c)?; + Ok(Self { + aspp1, + aspp_deforms_0, + aspp_deforms_1, + aspp_deforms_2, + global_avg_pool_1, + global_avg_pool_2, + conv1, + bn1, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let x1 = self.aspp1.forward(xs)?; + let x_aspp_deforms_0 = self.aspp_deforms_0.forward(xs)?; + let x_aspp_deforms_1 = self.aspp_deforms_1.forward(xs)?; + let x_aspp_deforms_2 = self.aspp_deforms_2.forward(xs)?; + + let (_, _, h, w) = xs.dims4()?; + assert_eq!(h, w, "avg_pool2d h, w mus be equal"); + let x5 = xs.avg_pool2d(h)?; + let x5 = self.global_avg_pool_1.forward(&x5)?; + let x5 = self.global_avg_pool_2.forward_t(&x5, false)?.relu()?; + let (_, _, h, w) = x1.dims4()?; + let x5 = interpolate_bilinear(&x5, (h, w), Some(true))?; + let xs = Tensor::cat( + &[x1, x_aspp_deforms_0, x_aspp_deforms_1, x_aspp_deforms_2, x5], + 1, + )?; + let xs = self.conv1.forward(&xs)?; + let xs = self.bn1.forward_t(&xs, false)?.relu()?; + Ok(xs) + } +} + +struct BasicDecBlk { + conv_in: Conv2d, + dec_att: ASPPDeformable, + conv_out: Conv2d, + bn_in: BatchNorm, + bn_out: BatchNorm, +} + +impl BasicDecBlk { + pub fn new(vb: VarBuilder, in_c: usize, out_c: usize) -> Result { + let inter_channels = 64; + let conv_in = get_conv2d(vb.pp("conv_in"), in_c, inter_channels, 3, 1, 1, 1, 1, true)?; + let dec_att = ASPPDeformable::new(vb.pp("dec_att"), inter_channels, inter_channels, vec![ + 1, 3, 7, + ])?; + let conv_out = get_conv2d( + vb.pp("conv_out"), + inter_channels, + out_c, + 3, + 1, + 1, + 1, + 1, + true, + )?; + let bn_in = get_batch_norm(vb.pp("bn_in"), 1e-5, inter_channels)?; + let bn_out = get_batch_norm(vb.pp("bn_out"), 1e-5, out_c)?; + Ok(Self { + conv_in, + dec_att, + conv_out, + bn_in, + bn_out, + }) + } + pub fn forward(&self, xs: &Tensor) -> Result { + let xs = self.conv_in.forward(xs)?; + let xs = self.bn_in.forward_t(&xs, false)?; + let xs = xs.relu()?; + let xs = self.dec_att.forward(&xs)?; + let xs = self.conv_out.forward(&xs)?; + let xs = self.bn_out.forward_t(&xs, false)?; + Ok(xs) + } +} + +struct SimpleConvs { + conv1: Conv2d, + conv_out: Conv2d, +} + +impl SimpleConvs { + pub fn new(vb: VarBuilder, in_c: usize, out_c: usize, inter_c: usize) -> Result { + // inter_c = 64 + let conv1 = get_conv2d(vb.pp("conv1"), in_c, inter_c, 3, 1, 1, 1, 1, true)?; + let conv_out = get_conv2d(vb.pp("conv_out"), inter_c, out_c, 3, 1, 1, 1, 1, true)?; + Ok(Self { conv1, conv_out }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + let x = self.conv1.forward(x)?; + let x = self.conv_out.forward(&x)?; + Ok(x) + } +} + +struct Conv2dWithBN { + conv_0: Conv2d, + bn_1: BatchNorm, +} + +impl Conv2dWithBN { + pub fn new( + vb: VarBuilder, + in_c: usize, + out_c: usize, + ks: usize, + padding: usize, + stride: usize, + ) -> Result { + let conv_0 = get_conv2d(vb.pp("0"), in_c, out_c, ks, padding, stride, 1, 1, true)?; + let bn_1 = get_batch_norm(vb.pp("1"), 1e-5, out_c)?; + Ok(Self { conv_0, bn_1 }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + let x = self.conv_0.forward(x)?; + let x = self.bn_1.forward_t(&x, false)?.relu()?; + Ok(x) + } +} + +struct Decoder { + ipt_blk5: SimpleConvs, + ipt_blk4: SimpleConvs, + ipt_blk3: SimpleConvs, + ipt_blk2: SimpleConvs, + ipt_blk1: SimpleConvs, + decoder_block4: BasicDecBlk, + decoder_block3: BasicDecBlk, + decoder_block2: BasicDecBlk, + decoder_block1: BasicDecBlk, + conv_out1: Conv2d, + // BasicLatBlk : conv, kernel_size=1 + lateral_block4: Conv2d, + lateral_block3: Conv2d, + lateral_block2: Conv2d, + // conv_ms_spvn_4: Conv2d, + // conv_ms_spvn_3: Conv2d, + // conv_ms_spvn_2: Conv2d, + gdt_convs_4: Conv2dWithBN, + gdt_convs_3: Conv2dWithBN, + gdt_convs_2: Conv2dWithBN, + gdt_convs_attn_4: Conv2d, + gdt_convs_attn_3: Conv2d, + gdt_convs_attn_2: Conv2d, +} + +impl Decoder { + pub fn new(vb: VarBuilder, channels: Vec) -> Result { + let ic = 64; + let ipt_blk5 = + SimpleConvs::new(vb.pp("ipt_blk5"), 2usize.pow(10) * 3, channels[0] / 8, ic)?; + let ipt_blk4 = SimpleConvs::new(vb.pp("ipt_blk4"), 2usize.pow(8) * 3, channels[0] / 8, ic)?; + let ipt_blk3 = SimpleConvs::new(vb.pp("ipt_blk3"), 2usize.pow(6) * 3, channels[1] / 8, ic)?; + let ipt_blk2 = SimpleConvs::new(vb.pp("ipt_blk2"), 2usize.pow(4) * 3, channels[2] / 8, ic)?; + let ipt_blk1 = SimpleConvs::new(vb.pp("ipt_blk1"), 3, channels[3] / 8, ic)?; + + let decoder_block4 = BasicDecBlk::new( + vb.pp("decoder_block4"), + channels[0] + channels[0] / 8, + channels[1], + )?; + let decoder_block3 = BasicDecBlk::new( + vb.pp("decoder_block3"), + channels[1] + channels[0] / 8, + channels[2], + )?; + let decoder_block2 = BasicDecBlk::new( + vb.pp("decoder_block2"), + channels[2] + channels[1] / 8, + channels[3], + )?; + let decoder_block1 = BasicDecBlk::new( + vb.pp("decoder_block1"), + channels[3] + channels[2] / 8, + channels[3] / 2, + )?; + + let conv_out1 = get_conv2d( + vb.pp("conv_out1.0"), + channels[3] / 2 + channels[3] / 8, + 1, + 1, + 0, + 1, + 1, + 1, + true, + )?; + let lateral_block4 = get_conv2d( + vb.pp("lateral_block4.conv"), + channels[1], + channels[1], + 1, + 0, + 1, + 1, + 1, + true, + )?; + let lateral_block3 = get_conv2d( + vb.pp("lateral_block3.conv"), + channels[2], + channels[2], + 1, + 0, + 1, + 1, + 1, + true, + )?; + let lateral_block2 = get_conv2d( + vb.pp("lateral_block2.conv"), + channels[3], + channels[3], + 1, + 0, + 1, + 1, + 1, + true, + )?; + + // let conv_ms_spvn_4 = + // get_conv2d(vb.pp("conv_ms_spvn_4"), channels[1], 1, 1, 0, 1, 1, 1, true)?; + // let conv_ms_spvn_3 = + // get_conv2d(vb.pp("conv_ms_spvn_3"), channels[2], 1, 1, 0, 1, 1, 1, true)?; + // let conv_ms_spvn_2 = + // get_conv2d(vb.pp("conv_ms_spvn_2"), channels[3], 1, 1, 0, 1, 1, 1, true)?; + let n = 16usize; + let gdt_convs_4 = Conv2dWithBN::new(vb.pp("gdt_convs_4"), channels[1], n, 3, 1, 1)?; + let gdt_convs_3 = Conv2dWithBN::new(vb.pp("gdt_convs_3"), channels[2], n, 3, 1, 1)?; + let gdt_convs_2 = Conv2dWithBN::new(vb.pp("gdt_convs_2"), channels[3], n, 3, 1, 1)?; + + let gdt_convs_attn_4 = get_conv2d(vb.pp("gdt_convs_attn_4.0"), n, 1, 1, 0, 1, 1, 1, true)?; + let gdt_convs_attn_3 = get_conv2d(vb.pp("gdt_convs_attn_3.0"), n, 1, 1, 0, 1, 1, 1, true)?; + let gdt_convs_attn_2 = get_conv2d(vb.pp("gdt_convs_attn_2.0"), n, 1, 1, 0, 1, 1, 1, true)?; + Ok(Self { + ipt_blk5, + ipt_blk4, + ipt_blk3, + ipt_blk2, + ipt_blk1, + decoder_block4, + decoder_block3, + decoder_block2, + decoder_block1, + conv_out1, + lateral_block4, + lateral_block3, + lateral_block2, + // conv_ms_spvn_4, + // conv_ms_spvn_3, + // conv_ms_spvn_2, + gdt_convs_4, + gdt_convs_3, + gdt_convs_2, + gdt_convs_attn_4, + gdt_convs_attn_3, + gdt_convs_attn_2, + }) + } + + pub fn get_patches_batch(&self, x: &Tensor, p: &Tensor) -> Result { + let (_, _, h, w) = p.dims4()?; + let mut patches_batch = vec![]; + for idx in 0..x.dim(0)? { + let x_i = x.i(idx)?.unsqueeze(0)?; // 保证bs的维度存在 + let columns_x = split_tensor_with_size(&x_i, w, D::Minus1)?; + let mut patches_x = vec![]; + for col_x in columns_x { + let pat_x = split_tensor_with_size(&col_x, h, D::Minus2)?; + patches_x.extend_from_slice(&pat_x); + } + let patch_sample = Tensor::cat(&patches_x, 1)?; + patches_batch.push(patch_sample); + } + let patch = Tensor::cat(&patches_batch, 0)?; + Ok(patch) + } + + pub fn forward(&self, features: Vec<&Tensor>) -> Result { + let [x, x1, x2, x3, x4] = features[..] else { + return Err(anyhow!(format!( + "swintransformer output exactly 3 elements" + ))); + }; + // let mut outs = vec![]; + let patches_batch = self.get_patches_batch(x, x4)?; + let (_, _, x4_h, x4_w) = x4.dims4()?; + let patches_batch = interpolate_bilinear(&patches_batch, (x4_h, x4_w), Some(true))?; + let ipt_blk5_out = self.ipt_blk5.forward(&patches_batch)?; + let x4 = Tensor::cat(&[x4, &ipt_blk5_out], 1)?; + let p4 = self.decoder_block4.forward(&x4)?; + // let m4 = self.conv_ms_spvn_4.forward(&p4)?; + // outs.push(m4); + let p4_gdt = self.gdt_convs_4.forward(&p4)?; + let gdt_attn_4 = sigmoid(&self.gdt_convs_attn_4.forward(&p4_gdt)?)?; + let p4 = p4.broadcast_mul(&gdt_attn_4)?; + + let (_, _, x3_h, x3_w) = x3.dims4()?; + let p4_inter = interpolate_bilinear(&p4, (x3_h, x3_w), Some(true))?; + let p3_ = self.lateral_block4.forward(x3)?; + let p3_ = p4_inter.add(&p3_)?; + let patches_batch = self.get_patches_batch(x, &p3_)?; + let patches_batch = interpolate_bilinear(&patches_batch, (x3_h, x3_w), Some(true))?; + let ipt_blk4_out = self.ipt_blk4.forward(&patches_batch)?; + let p3_ = Tensor::cat(&[p3_, ipt_blk4_out], 1)?; + let p3 = self.decoder_block3.forward(&p3_)?; + // let m3 = self.conv_ms_spvn_3.forward(&p3)?; + // outs.push(m3); + let p3_gdt = self.gdt_convs_3.forward(&p3)?; + let gdt_attn_3 = sigmoid(&self.gdt_convs_attn_3.forward(&p3_gdt)?)?; + let p3 = p3.broadcast_mul(&gdt_attn_3)?; + + let (_, _, x2_h, x2_w) = x2.dims4()?; + let p3_inter = interpolate_bilinear(&p3, (x2_h, x2_w), Some(true))?; + let p2_ = self.lateral_block3.forward(x2)?; + let p2_ = p3_inter.add(&p2_)?; + let patches_batch = self.get_patches_batch(x, &p2_)?; + let patches_batch = interpolate_bilinear(&patches_batch, (x2_h, x2_w), Some(true))?; + let ipt_blk3_out = self.ipt_blk3.forward(&patches_batch)?; + let p2_ = Tensor::cat(&[p2_, ipt_blk3_out], 1)?; + let p2 = self.decoder_block2.forward(&p2_)?; + // let m2 = self.conv_ms_spvn_2.forward(&p2)?; + // outs.push(m2); + let p2_gdt = self.gdt_convs_2.forward(&p2)?; + let gdt_attn_2 = sigmoid(&self.gdt_convs_attn_2.forward(&p2_gdt)?)?; + let p2 = p2.broadcast_mul(&gdt_attn_2)?; + + let (_, _, x1_h, x1_w) = x1.dims4()?; + let p2_inter = interpolate_bilinear(&p2, (x1_h, x1_w), Some(true))?; + let p1_ = self.lateral_block2.forward(x1)?; + let p1_ = p2_inter.add(&p1_)?; + let patches_batch = self.get_patches_batch(x, &p1_)?; + let patches_batch = interpolate_bilinear(&patches_batch, (x1_h, x1_w), Some(true))?; + let ipt_blk2_out = self.ipt_blk2.forward(&patches_batch)?; + let p1_ = Tensor::cat(&[p1_, ipt_blk2_out], 1)?; + let p1_ = self.decoder_block1.forward(&p1_)?; + + let (_, _, x_h, x_w) = x.dims4()?; + let p1_ = interpolate_bilinear(&p1_, (x_h, x_w), Some(true))?; + // let patches_batch = self.get_patches_batch(x, &p1_)?; + // let patches_batch = interpolate_bilinear(&patches_batch, (x_h, x_w), Some(true))?; + let ipt_blk1_out = self.ipt_blk1.forward(x)?; + let p1_ = Tensor::cat(&[p1_, ipt_blk1_out], 1)?; + let p1_out = self.conv_out1.forward(&p1_)?; + let out = sigmoid(&p1_out)?; + // outs.push(p1_out); + Ok(out) + } +} + +pub struct BiRefNet { + bb: SwinTransformer, + squeeze_module_0: BasicDecBlk, + decoder: Decoder, +} +impl BiRefNet { + pub fn new(vb: VarBuilder) -> Result { + let bb = SwinTransformer::new( + vb.pp("bb"), + 4, + 3, + 192, + vec![2, 2, 18, 2], + vec![6, 12, 24, 48], + 12, + 4.0, + true, + true, + vec![0, 1, 2, 3], + )?; + let channels = vec![3072, 1536, 768, 384]; + let in_c = channels.iter().sum(); + let squeeze_module_0 = BasicDecBlk::new(vb.pp("squeeze_module.0"), in_c, channels[0])?; + let decoder = Decoder::new(vb.pp("decoder"), channels)?; + Ok(Self { + bb, + squeeze_module_0, + decoder, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let [ref x1, ref x2, ref x3, ref x4] = self.bb.forward(xs)?[..] else { + return Err(anyhow!(format!( + "swintransformer output exactly 3 elements" + ))); + }; + + let (_, _, h, w) = xs.dims4()?; + let cat_xs = interpolate_bilinear(xs, (h / 2, w / 2), Some(true))?; + let [ref x1_, ref x2_, ref x3_, ref x4_] = self.bb.forward(&cat_xs)?[..] else { + return Err(anyhow!(format!( + "swintransformer output exactly 3 elements" + ))); + }; + let (_, _, x1_h, x1_w) = x1.dims4()?; + let x1_ = interpolate_bilinear(x1_, (x1_h, x1_w), Some(true))?; + let x1 = Tensor::cat(&[x1, &x1_], 1)?; + let (_, _, x2_h, x2_w) = x2.dims4()?; + let x2_ = interpolate_bilinear(x2_, (x2_h, x2_w), Some(true))?; + let x2 = Tensor::cat(&[x2, &x2_], 1)?; + let (_, _, x3_h, x3_w) = x3.dims4()?; + let x3_ = interpolate_bilinear(x3_, (x3_h, x3_w), Some(true))?; + let x3 = Tensor::cat(&[x3, &x3_], 1)?; + let (_, _, x4_h, x4_w) = x4.dims4()?; + let x4_ = interpolate_bilinear(x4_, (x4_h, x4_w), Some(true))?; + let x4 = Tensor::cat(&[x4, &x4_], 1)?; + let x1_resize = interpolate_bilinear(&x1, (x4_h, x4_w), Some(true))?; + let x2_resize = interpolate_bilinear(&x2, (x4_h, x4_w), Some(true))?; + let x3_resize = interpolate_bilinear(&x3, (x4_h, x4_w), Some(true))?; + let x4 = Tensor::cat(&[x1_resize, x2_resize, x3_resize, x4], 1)?; + + let x4 = self.squeeze_module_0.forward(&x4)?; + + let features = vec![xs, &x1, &x2, &x3, &x4]; + let output = self.decoder.forward(features)?; + Ok(output) + } } diff --git a/src/models/rmbg2_0/processor.rs b/src/models/rmbg2_0/processor.rs deleted file mode 100644 index e69de29..0000000 diff --git a/src/utils/img_utils.rs b/src/utils/img_utils.rs index 99ff264..1dba19b 100644 --- a/src/utils/img_utils.rs +++ b/src/utils/img_utils.rs @@ -285,3 +285,40 @@ pub fn img_smart_resize( } Ok((h_bar, w_bar)) } + +pub fn img_transform_with_resize( + img: &DynamicImage, + h: u32, + w: u32, + mean: &Tensor, + std: &Tensor, + device: &Device, + dtype: DType, +) -> Result { + let img_resize = img.resize_exact(w, h, imageops::FilterType::CatmullRom); + let img_tensor = img_transform(&img_resize, mean, std, device, dtype)?; + Ok(img_tensor) +} + +pub fn float_tensor_to_dynamic_image(tensor: &Tensor) -> Result { + let tensor = tensor.affine(255.0, 0.0)?.clamp(0.0, 255.0)?; + let tensor_u8 = tensor.to_dtype(DType::U8)?.to_device(&Device::Cpu)?; + let (c, h, w) = tensor_u8.dims3()?; + match c { + 1 => { + let tensor_u8 = tensor_u8.reshape((h, w))?; + let data: Vec = tensor_u8.flatten_all()?.to_vec1()?; + let img = ImageBuffer::from_raw(w as u32, h as u32, data) + .ok_or_else(|| anyhow!("Failed to create image buffer"))?; + Ok(DynamicImage::ImageLuma8(img)) + } + 3 => { + let tensor_u8 = tensor_u8.permute((1, 2, 0))?; + let data: Vec = tensor_u8.flatten_all()?.to_vec1()?; + let img = ImageBuffer::from_raw(w as u32, h as u32, data) + .ok_or_else(|| anyhow!("Failed to create image buffer"))?; + Ok(DynamicImage::ImageRgb8(img)) + } + _ => Err(anyhow!(format!("Unsupported number of channels: {}", c))), + } +} diff --git a/src/utils/tensor_utils.rs b/src/utils/tensor_utils.rs index 1057147..fbee88d 100644 --- a/src/utils/tensor_utils.rs +++ b/src/utils/tensor_utils.rs @@ -51,6 +51,9 @@ pub fn repeat_kv(xs: Tensor, n_rep: usize) -> Result { } pub fn split_tensor(t: &Tensor, splits: &[usize], dim: D) -> Result> { + // 按给定长度切分tensor + // 例: t:(25), splits: [5, 10, 5, 5] dim: 0, + // 返回vec len=4, 其中tensor维度分别是:(5), (10), (5), (5) let dim = dim.to_index(t.shape(), "split")?; let mut split_res = Vec::new(); let mut index = 0; @@ -61,6 +64,28 @@ pub fn split_tensor(t: &Tensor, splits: &[usize], dim: D) -> Result( + t: &Tensor, + splits_size: usize, + dim: D, +) -> Result> { + // 按给定size切分tensor + // 例: t:(25), splits: 5 dim: 0, + // 返回vec len=5, 其中tensor维度分别是:(5), (5), (5), (5), (5) + let dim = dim.to_index(t.shape(), "split")?; + let mut split_res = Vec::new(); + let dim_size = t.dim(dim)?; + assert_eq!( + dim_size % splits_size, + 0, + "input tensor dim size % splits_size must be equal to 0" + ); + for split in (0..dim_size).step_by(splits_size) { + split_res.push(t.narrow(dim, split, splits_size)?); + } + Ok(split_res) +} + pub fn safe_arg_sort_last_dim(t: &Tensor, ascending: bool) -> Result { // tensor在GPU上时,维度超过1024, arg_sort_last_dim方法会报错 // 所以维度大于1024时,放到CPU上处理 @@ -230,7 +255,8 @@ pub fn get_not_equal_mask(input_ids: &Tensor, token_ids: u32) -> Result } pub fn get_equal_mask(input_ids: &Tensor, token_ids: u32) -> Result { - let image_token_id_tensor = Tensor::new(vec![token_ids], input_ids.device())?; + let image_token_id_tensor = + Tensor::new(vec![token_ids], input_ids.device())?.to_dtype(input_ids.dtype())?; let mask = input_ids .broadcast_eq(&image_token_id_tensor)? .to_dtype(candle_core::DType::U32)?; diff --git a/tests/messy_test.rs b/tests/messy_test.rs index 9626571..b0006f5 100644 --- a/tests/messy_test.rs +++ b/tests/messy_test.rs @@ -1,22 +1,62 @@ -use std::{path::PathBuf, str::FromStr}; - use anyhow::Result; +use candle_core::Tensor; #[test] fn messy_test() -> Result<()> { // RUST_BACKTRACE=1 cargo test -F cuda messy_test -r -- --nocapture - let path_str = "file://./assets/img/ocr_test1.png"; - let path = url::Url::from_str(path_str)?; - let path = path.to_file_path(); - let path = match path { - Ok(path) => path, - Err(_) => { - let mut path = path_str.to_owned(); - path = path.split_off(7); - PathBuf::from(path) - } - }; - println!("to file path: {:?}", path); + let device = &candle_core::Device::Cpu; + let x = Tensor::arange(0.0, 9.0, device)?; + println!("x: {}", x); + let x = x + .unsqueeze(0)? + .unsqueeze(0)? + .broadcast_as((5, 5, 9))? + .reshape((5, 5, 3, 3))?; + println!("x: {}", x); + let x = x.permute((0, 2, 1, 3))?; + println!("x: {}", x); + let x = x.reshape((15, 15))?; + println!("x: {}", x); + // let xs = Tensor::rand(0.0, 5.0, (1, 1, 3, 3), device)?; + // println!("xs: {}", xs); + // let xs = xs.pad_with_zeros(3, 2, 2)? + // .pad_with_zeros(2, 2, 2)?; + // println!("xs: {}", xs); + // let xs = Tensor::arange(0.0, 25.0, device)?; + // println!("xs: {}", xs); + // let splits = split_tensor_with_size(&xs, 5, 0)?; + // for v in splits { + // println!("v: {}", v); + // } + // let xs = Tensor::arange(0.0, 25.0, device)?.broadcast_as((1, 1, 5, 5))?; + // println!("xs: {}", xs); + // let xs = xs.avg_pool2d(5)?; + // println!("xs: {}", xs); + // let xs = Tensor::rand(0.0, 1.0, (1, 4, 4, 2), device)?; + // println!("xs: {}", xs); + // let shape = Shape::from_dims(&[1, 2, 2, 2, 2, 2]); + // let xs = xs.reshape(shape)?; + // println!("xs: {}", xs); + // let x0 = xs.i((.., .., 0, .., 0, ..))?; + // let x1 = xs.i((.., .., 1, .., 0, ..))?; + // let x2 = xs.i((.., .., 0, .., 1, ..))?; + // let x3 = xs.i((.., .., 1, .., 1, ..))?; + // let xs = Tensor::cat(&[x0, x1, x2, x3], D::Minus1)?; + // println!("xs: {}", xs); + // let xs = xs.reshape((1, (), 4 * 2))?; + // println!("xs: {}", xs); + // let path_str = "file://./assets/img/ocr_test1.png"; + // let path = url::Url::from_str(path_str)?; + // let path = path.to_file_path(); + // let path = match path { + // Ok(path) => path, + // Err(_) => { + // let mut path = path_str.to_owned(); + // path = path.split_off(7); + // PathBuf::from(path) + // } + // }; + // println!("to file path: {:?}", path); // let device = &candle_core::Device::Cpu; // let t = Tensor::arange(0.0f32, 40.0, device)?.broadcast_as((1, 1, 40, 40))?; diff --git a/tests/test_rmbg2_0.rs b/tests/test_rmbg2_0.rs new file mode 100644 index 0000000..463c4ef --- /dev/null +++ b/tests/test_rmbg2_0.rs @@ -0,0 +1,47 @@ +use std::time::Instant; + +use aha::models::rmbg2_0::generate::RMBG2_0; +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use anyhow::Result; + +#[test] +fn rmbg2_0_generate() -> Result<()> { + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda rmbg2_0_generate -r -- --nocapture + + let model_path = "/home/jhq/huggingface_model/AI-ModelScope/RMBG-2.0/"; + + let message = r#" + { + "model": "rmbg2.0", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "image", + "image_url": + { + "url": "file://./assets/img/gougou.jpg" + } + } + ] + } + ] + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let model = RMBG2_0::init(model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + + let i_start = Instant::now(); + let result = model.generate(mes)?; + for (i, img) in result.iter().enumerate() { + let _ = img.save(format!("rmbg_{i}.png")); + } + let i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + + Ok(()) +} diff --git a/tests/test_voxcpm.rs b/tests/test_voxcpm.rs index 290e9b1..dddd2d1 100644 --- a/tests/test_voxcpm.rs +++ b/tests/test_voxcpm.rs @@ -19,7 +19,7 @@ fn voxcpm_generate() -> Result<()> { let i_start = Instant::now(); // let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?; let generate = voxcpm_generate.generate( - "太阳当空照,花儿对我笑,小鸟说早早早".to_string(), + "VoxCPM is an innovative end-to-end TTS model from ModelBest, designed to generate highly realistic speech.".to_string(), Some("啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string()), Some("./assets/audio/voice_01.wav".to_string()), // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()), diff --git a/tests/test_voxcpm1_5.rs b/tests/test_voxcpm1_5.rs index 29332f9..cf2dca7 100644 --- a/tests/test_voxcpm1_5.rs +++ b/tests/test_voxcpm1_5.rs @@ -19,7 +19,7 @@ fn voxcpm1_5_generate() -> Result<()> { let i_start = Instant::now(); // let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?; let generate = voxcpm_generate.generate( - "太阳当空照,花儿对我笑,小鸟说早早早".to_string(), + "VoxCPM is an innovative end-to-end TTS model from ModelBest, designed to generate highly realistic speech.".to_string(), Some("啥子小师叔,打狗还要看主人,你再要继续,我就是你的对手".to_string()), Some("./assets/audio/voice_01.wav".to_string()), // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),