add RMBGv2.0

This commit is contained in:
jhqxxx
2025-12-23 19:23:21 +08:00
parent 675fd6f89e
commit 06bd7fce04
16 changed files with 1576 additions and 214 deletions
+27 -1
View File
@@ -51,6 +51,9 @@ pub fn repeat_kv(xs: Tensor, n_rep: usize) -> Result<Tensor> {
}
pub fn split_tensor<D: Dim>(t: &Tensor, splits: &[usize], dim: D) -> Result<Vec<Tensor>> {
// 按给定长度切分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<D: Dim>(t: &Tensor, splits: &[usize], dim: D) -> Result<Vec<
Ok(split_res)
}
pub fn split_tensor_with_size<D: Dim>(
t: &Tensor,
splits_size: usize,
dim: D,
) -> Result<Vec<Tensor>> {
// 按给定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> {
// 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<Tensor>
}
pub fn get_equal_mask(input_ids: &Tensor, token_ids: u32) -> Result<Tensor> {
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)?;