update code format

This commit is contained in:
jhqxxx
2025-10-26 21:57:53 +08:00
parent 6244ffe324
commit 0e4ca66aca
12 changed files with 54 additions and 52 deletions
+12 -18
View File
@@ -247,7 +247,7 @@ pub fn bitor_tensor(mask1: &Tensor, mask2: &Tensor) -> Result<Tensor> {
mask1.shape() == mask2.shape(),
" bitor_tensor two tensor shape mask be equal"
);
let bitor = mask1.add(&mask2)?.ne(&Tensor::zeros_like(&mask1)?)?;
let bitor = mask1.add(mask2)?.ne(&Tensor::zeros_like(mask1)?)?;
Ok(bitor)
}
@@ -256,38 +256,37 @@ pub fn prod_tensor_last_dim(t: &Tensor) -> Result<Tensor> {
0 => t.clone(),
1 => {
let data_type = t.dtype();
let t_prod = match data_type {
match data_type {
DType::U8 => {
let t_vec = t.to_vec1::<u8>()?;
let prod = t_vec.iter().product::<u8>();
Tensor::from_slice(&vec![prod], 1, t.device())?
Tensor::from_slice(&[prod], 1, t.device())?
}
DType::U32 => {
let t_vec = t.to_vec1::<u32>()?;
let prod = t_vec.iter().product::<u32>();
Tensor::from_slice(&vec![prod], 1, t.device())?
Tensor::from_slice(&[prod], 1, t.device())?
}
DType::I64 => {
let t_vec = t.to_vec1::<i64>()?;
let prod = t_vec.iter().product::<i64>();
Tensor::from_slice(&vec![prod], 1, t.device())?
Tensor::from_slice(&[prod], 1, t.device())?
}
DType::F64 => {
let t_vec = t.to_vec1::<f64>()?;
let prod = t_vec.iter().product::<f64>();
Tensor::from_slice(&vec![prod], 1, t.device())?
Tensor::from_slice(&[prod], 1, t.device())?
}
_ => {
let t_vec = t.to_vec1::<f32>()?;
let prod = t_vec.iter().product::<f32>();
Tensor::from_slice(&vec![prod], 1, t.device())?
Tensor::from_slice(&[prod], 1, t.device())?
}
};
t_prod
}
}
2 => {
let data_type = t.dtype();
let t_prod = match data_type {
match data_type {
DType::U8 => {
let t_vec = t.to_vec2::<u8>()?;
let mut prod_vec = vec![];
@@ -333,8 +332,7 @@ pub fn prod_tensor_last_dim(t: &Tensor) -> Result<Tensor> {
}
Tensor::new(prod_vec, t.device())?
}
};
t_prod
}
}
_ => {
return Err(anyhow!(format!("can not action this dim")));
@@ -343,12 +341,8 @@ pub fn prod_tensor_last_dim(t: &Tensor) -> Result<Tensor> {
Ok(prod)
}
pub fn mask_index_add(
original: &Tensor,
mask: &Tensor,
add: &Tensor,
) -> Result<Tensor> {
let visual_nonzero_index = nonzero_index(&mask)?;
pub fn mask_index_add(original: &Tensor, mask: &Tensor, add: &Tensor) -> Result<Tensor> {
let visual_nonzero_index = nonzero_index(mask)?;
let xs = original.index_add(&visual_nonzero_index, add, 0)?;
Ok(xs)
}