2025-10-15 21:03:49 +08:00
|
|
|
|
use anyhow::{Ok, Result, anyhow};
|
2025-09-22 16:38:12 +08:00
|
|
|
|
use candle_core::{D, DType, Device, IndexOp, Tensor, shape::Dim};
|
2025-11-13 00:33:27 +08:00
|
|
|
|
use rocket::figment::value;
|
2025-09-22 16:38:12 +08:00
|
|
|
|
|
2025-09-25 12:09:25 +08:00
|
|
|
|
pub fn prepare_causal_attention_mask(
|
|
|
|
|
|
b_size: usize,
|
|
|
|
|
|
tgt_len: usize,
|
|
|
|
|
|
seqlen_offset: usize,
|
2025-10-15 21:03:49 +08:00
|
|
|
|
device: &Device,
|
2025-09-25 12:09:25 +08:00
|
|
|
|
) -> Result<Tensor> {
|
|
|
|
|
|
// Sliding window mask?
|
2025-11-13 00:33:27 +08:00
|
|
|
|
// 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)?)?;
|
2025-09-25 12:09:25 +08:00
|
|
|
|
let mask = if seqlen_offset > 0 {
|
2025-10-11 13:14:51 +08:00
|
|
|
|
let mask0 = Tensor::zeros((tgt_len, seqlen_offset), DType::F32, device)?;
|
2025-09-25 12:09:25 +08:00
|
|
|
|
Tensor::cat(&[&mask0, &mask], D::Minus1)?
|
|
|
|
|
|
} else {
|
|
|
|
|
|
mask
|
|
|
|
|
|
};
|
|
|
|
|
|
let mask = mask
|
|
|
|
|
|
.expand((b_size, 1, tgt_len, tgt_len + seqlen_offset))?
|
2025-10-11 13:14:51 +08:00
|
|
|
|
.to_dtype(DType::F32)?;
|
2025-09-25 12:09:25 +08:00
|
|
|
|
Ok(mask)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
pub fn repeat_kv(xs: Tensor, n_rep: usize) -> Result<Tensor> {
|
|
|
|
|
|
if n_rep == 1 {
|
|
|
|
|
|
Ok(xs)
|
|
|
|
|
|
} else {
|
|
|
|
|
|
let (b_sz, n_kv_head, seq_len, head_dim) = xs.dims4()?;
|
|
|
|
|
|
// Using cat is faster than a broadcast as it avoids going through a potentially
|
|
|
|
|
|
// strided copy.
|
|
|
|
|
|
// https://github.com/huggingface/candle/pull/2043
|
|
|
|
|
|
let kv = Tensor::cat(&vec![&xs; n_rep], 2)?.reshape((
|
|
|
|
|
|
b_sz,
|
|
|
|
|
|
n_kv_head * n_rep,
|
|
|
|
|
|
seq_len,
|
|
|
|
|
|
head_dim,
|
|
|
|
|
|
))?;
|
|
|
|
|
|
Ok(kv)
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2025-10-26 21:39:23 +08:00
|
|
|
|
pub fn split_tensor<D: Dim>(t: &Tensor, splits: &[usize], dim: D) -> Result<Vec<Tensor>> {
|
2025-09-22 16:38:12 +08:00
|
|
|
|
let dim = dim.to_index(t.shape(), "split")?;
|
|
|
|
|
|
let mut split_res = Vec::new();
|
|
|
|
|
|
let mut index = 0;
|
|
|
|
|
|
for split in splits {
|
|
|
|
|
|
split_res.push(t.narrow(dim, index, *split)?);
|
|
|
|
|
|
index += *split;
|
|
|
|
|
|
}
|
|
|
|
|
|
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上处理
|
|
|
|
|
|
let last_dim = t.dims()[t.rank() - 1];
|
|
|
|
|
|
if last_dim <= 1024 {
|
|
|
|
|
|
let t = t.arg_sort_last_dim(ascending)?;
|
|
|
|
|
|
Ok(t)
|
|
|
|
|
|
} else {
|
|
|
|
|
|
let cpu_tensor = t.to_device(&Device::Cpu)?;
|
|
|
|
|
|
let sorted_indices = cpu_tensor.arg_sort_last_dim(ascending)?;
|
|
|
|
|
|
let t = sorted_indices.to_device(t.device())?;
|
|
|
|
|
|
Ok(t)
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
pub fn nonzero_index_vec(mask: &Tensor) -> Result<Vec<u32>> {
|
|
|
|
|
|
// 根据mask矩阵选出其中不为0的元素所在索引, 返回vec
|
|
|
|
|
|
// 只能处理1维数据
|
|
|
|
|
|
let mut mask = mask.clone();
|
|
|
|
|
|
if mask.dtype() != DType::U32 {
|
|
|
|
|
|
mask = mask.to_dtype(DType::U32)?;
|
|
|
|
|
|
}
|
|
|
|
|
|
match mask.rank() {
|
2025-10-15 21:03:49 +08:00
|
|
|
|
0 => Err(anyhow!(format!(
|
|
|
|
|
|
"input rank must > 0, the input tensor rank: {}",
|
|
|
|
|
|
mask.rank()
|
|
|
|
|
|
))),
|
2025-09-22 16:38:12 +08:00
|
|
|
|
1 => {
|
|
|
|
|
|
let mask_vector = mask.to_vec1::<u32>()?;
|
|
|
|
|
|
let indices: Vec<u32> = mask_vector
|
|
|
|
|
|
.iter()
|
|
|
|
|
|
.enumerate()
|
|
|
|
|
|
.filter_map(|(idx, &val)| if val != 0 { Some(idx as u32) } else { None })
|
|
|
|
|
|
.collect();
|
|
|
|
|
|
Ok(indices)
|
|
|
|
|
|
}
|
2025-10-15 21:03:49 +08:00
|
|
|
|
_ => Err(anyhow!(format!(
|
|
|
|
|
|
"input rank not support, the input tensor rank: {}",
|
|
|
|
|
|
mask.rank()
|
|
|
|
|
|
))),
|
2025-09-22 16:38:12 +08:00
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
pub fn nonzero_index(mask: &Tensor) -> Result<Tensor> {
|
|
|
|
|
|
// 根据mask矩阵选出其中不为1的元素所在索引, 返回Tensor
|
|
|
|
|
|
let indices_tensor = match mask.rank() {
|
|
|
|
|
|
0 => {
|
|
|
|
|
|
return Err(anyhow!(format!(
|
|
|
|
|
|
"input rank must > 0, the input tensor rank: {}",
|
|
|
|
|
|
mask.rank()
|
|
|
|
|
|
)));
|
|
|
|
|
|
}
|
|
|
|
|
|
1 => {
|
|
|
|
|
|
let index_vec = nonzero_index_vec(mask)?;
|
2025-10-15 21:03:49 +08:00
|
|
|
|
Tensor::from_slice(&index_vec, index_vec.len(), mask.device())?
|
2025-09-22 16:38:12 +08:00
|
|
|
|
}
|
|
|
|
|
|
_ => {
|
|
|
|
|
|
return Err(anyhow!(format!(
|
|
|
|
|
|
"input rank must == 1, the input tensor rank: {}",
|
|
|
|
|
|
mask.rank()
|
|
|
|
|
|
)));
|
|
|
|
|
|
}
|
|
|
|
|
|
};
|
|
|
|
|
|
Ok(indices_tensor)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
pub fn zero_index_vec(mask: &Tensor) -> Result<Vec<u32>> {
|
|
|
|
|
|
// 根据mask矩阵选出其中为0的元素所在索引, 返回vec
|
|
|
|
|
|
// 只能处理1维数据
|
|
|
|
|
|
let mut mask = mask.clone();
|
|
|
|
|
|
if mask.dtype() != DType::U32 {
|
|
|
|
|
|
mask = mask.to_dtype(DType::U32)?;
|
|
|
|
|
|
}
|
|
|
|
|
|
match mask.rank() {
|
2025-10-15 21:03:49 +08:00
|
|
|
|
0 => Err(anyhow!(format!(
|
|
|
|
|
|
"input rank must > 0, the input tensor rank: {}",
|
|
|
|
|
|
mask.rank()
|
|
|
|
|
|
))),
|
2025-09-22 16:38:12 +08:00
|
|
|
|
1 => {
|
|
|
|
|
|
let mask_vector = mask.to_vec1::<u32>()?;
|
|
|
|
|
|
let indices: Vec<u32> = mask_vector
|
|
|
|
|
|
.iter()
|
|
|
|
|
|
.enumerate()
|
|
|
|
|
|
.filter_map(|(idx, &val)| if val == 0 { Some(idx as u32) } else { None })
|
|
|
|
|
|
.collect();
|
|
|
|
|
|
Ok(indices)
|
|
|
|
|
|
}
|
2025-10-15 21:03:49 +08:00
|
|
|
|
_ => Err(anyhow!(format!(
|
|
|
|
|
|
"input rank not support, the input tensor rank: {}",
|
|
|
|
|
|
mask.rank()
|
|
|
|
|
|
))),
|
2025-09-22 16:38:12 +08:00
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
pub fn zero_index(mask: &Tensor) -> Result<Tensor> {
|
|
|
|
|
|
let index_vec = zero_index_vec(mask)?;
|
|
|
|
|
|
let indices_tensor = Tensor::from_slice(&index_vec, index_vec.len(), mask.device())?;
|
|
|
|
|
|
Ok(indices_tensor)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
pub fn nonzero_slice(mask: &Tensor) -> Result<Vec<(usize, usize)>> {
|
|
|
|
|
|
// 根据mask矩阵选出其中非0的元素所在索引
|
|
|
|
|
|
// 根据索引获取连续索引间隔
|
|
|
|
|
|
// 如不为零索引元素为[0, 3, 4, 5, 8, 9]
|
|
|
|
|
|
// 间隔为: [(0, 1), (3, 6), (8, 10)]
|
|
|
|
|
|
// 索引前闭后开
|
|
|
|
|
|
let mut index_vec = nonzero_index_vec(mask)?;
|
|
|
|
|
|
match index_vec.len() {
|
2025-10-15 21:03:49 +08:00
|
|
|
|
0 => Ok(vec![]),
|
|
|
|
|
|
1 => Ok(vec![(index_vec[0] as usize, (index_vec[0] + 1) as usize)]),
|
2025-09-22 16:38:12 +08:00
|
|
|
|
_ => {
|
|
|
|
|
|
let mut vec_slice = vec![];
|
|
|
|
|
|
let mut start = index_vec.remove(0);
|
|
|
|
|
|
let mut last = start;
|
|
|
|
|
|
|
|
|
|
|
|
for i in index_vec {
|
|
|
|
|
|
if i == (last + 1) {
|
|
|
|
|
|
last = i;
|
|
|
|
|
|
continue;
|
|
|
|
|
|
} else {
|
|
|
|
|
|
vec_slice.push((start as usize, (last + 1) as usize));
|
|
|
|
|
|
start = i;
|
|
|
|
|
|
last = i;
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
vec_slice.push((start as usize, (last + 1) as usize));
|
|
|
|
|
|
Ok(vec_slice)
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
pub fn masked_scatter_dim0(original: &Tensor, replace: &Tensor, mask: &Tensor) -> Result<Tensor> {
|
|
|
|
|
|
// 根据mask中非0元素所在索引,使用replace中的数据替换掉original中的数据
|
|
|
|
|
|
// original: rank = 3: (bs, seq_len, hidden_dim)
|
|
|
|
|
|
// replace: rank = 2: (seq_len, hidden_dim)
|
|
|
|
|
|
// mask: rank = 2: (bs, seq_len)
|
|
|
|
|
|
// 推理时bs=1,为了方便替换,将bs squeeze,替换后再unsqueeze
|
|
|
|
|
|
// 按行替换
|
|
|
|
|
|
if original.dim(0)? != 1 || mask.dim(0)? != 1 {
|
|
|
|
|
|
return Err(anyhow!(format!(
|
|
|
|
|
|
"masked_scatter_dim0 original bs: {} or mask bs :{} not equal to 1 ",
|
|
|
|
|
|
original.dim(0)?,
|
|
|
|
|
|
mask.dim(0)? != 1
|
|
|
|
|
|
)));
|
|
|
|
|
|
}
|
|
|
|
|
|
let mut original = original.squeeze(0)?;
|
|
|
|
|
|
let mask = mask.squeeze(0)?;
|
|
|
|
|
|
let slices = nonzero_slice(&mask)?;
|
|
|
|
|
|
let mut sub_start = 0usize;
|
2025-10-10 20:36:52 +08:00
|
|
|
|
let mut sub_end;
|
2025-09-22 16:38:12 +08:00
|
|
|
|
for (start, end) in slices {
|
|
|
|
|
|
sub_end = sub_start + (end - start);
|
|
|
|
|
|
let sub_replace = replace.i((sub_start..sub_end, ..))?;
|
|
|
|
|
|
original = original.slice_assign(&[(start..end), (0..original.dim(1)?)], &sub_replace)?;
|
|
|
|
|
|
sub_start = sub_end;
|
|
|
|
|
|
}
|
|
|
|
|
|
original = original.unsqueeze(0)?;
|
|
|
|
|
|
Ok(original)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
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 mask = input_ids
|
|
|
|
|
|
.broadcast_eq(&image_token_id_tensor)?
|
|
|
|
|
|
.to_dtype(candle_core::DType::U32)?;
|
|
|
|
|
|
Ok(mask)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
pub fn get_vision_next_indices(input_ids: &Tensor, token_id: u32) -> Result<Tensor> {
|
|
|
|
|
|
// input_ids -> shape: (seq_len)
|
2025-10-15 21:03:49 +08:00
|
|
|
|
let mask = get_equal_mask(input_ids, token_id)?;
|
2025-09-22 16:38:12 +08:00
|
|
|
|
let indices = nonzero_index(&mask)?;
|
|
|
|
|
|
let indices = indices.broadcast_add(&Tensor::new(vec![1u32], input_ids.device())?)?;
|
|
|
|
|
|
Ok(indices)
|
|
|
|
|
|
}
|
2025-10-03 22:25:58 +08:00
|
|
|
|
|
|
|
|
|
|
pub fn linspace(start: f32, end: f32, steps: usize, device: &Device) -> Result<Tensor> {
|
|
|
|
|
|
assert!(steps > 0, "steps must be > 0");
|
|
|
|
|
|
if steps == 1 {
|
|
|
|
|
|
let t = Tensor::from_slice(&[start], 1, device)?;
|
|
|
|
|
|
return Ok(t);
|
2025-10-15 21:03:49 +08:00
|
|
|
|
}
|
|
|
|
|
|
let step_size = (end - start) / (steps - 1) as f32;
|
|
|
|
|
|
let data: Vec<f32> = (0..steps).map(|i| start + i as f32 * step_size).collect();
|
|
|
|
|
|
|
2025-10-03 22:25:58 +08:00
|
|
|
|
let t = Tensor::from_slice(&data, steps, device)?;
|
|
|
|
|
|
Ok(t)
|
2025-10-15 21:03:49 +08:00
|
|
|
|
}
|
2025-10-26 21:39:23 +08:00
|
|
|
|
|
|
|
|
|
|
pub fn bitor_tensor(mask1: &Tensor, mask2: &Tensor) -> Result<Tensor> {
|
|
|
|
|
|
assert!(
|
|
|
|
|
|
mask1.shape() == mask2.shape(),
|
|
|
|
|
|
" bitor_tensor two tensor shape mask be equal"
|
|
|
|
|
|
);
|
2025-10-26 21:57:53 +08:00
|
|
|
|
let bitor = mask1.add(mask2)?.ne(&Tensor::zeros_like(mask1)?)?;
|
2025-10-26 21:39:23 +08:00
|
|
|
|
Ok(bitor)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
pub fn prod_tensor_last_dim(t: &Tensor) -> Result<Tensor> {
|
|
|
|
|
|
let prod = match t.rank() {
|
|
|
|
|
|
0 => t.clone(),
|
|
|
|
|
|
1 => {
|
|
|
|
|
|
let data_type = t.dtype();
|
2025-10-26 21:57:53 +08:00
|
|
|
|
match data_type {
|
2025-10-26 21:39:23 +08:00
|
|
|
|
DType::U8 => {
|
|
|
|
|
|
let t_vec = t.to_vec1::<u8>()?;
|
|
|
|
|
|
let prod = t_vec.iter().product::<u8>();
|
2025-10-26 21:57:53 +08:00
|
|
|
|
Tensor::from_slice(&[prod], 1, t.device())?
|
2025-10-26 21:39:23 +08:00
|
|
|
|
}
|
|
|
|
|
|
DType::U32 => {
|
|
|
|
|
|
let t_vec = t.to_vec1::<u32>()?;
|
|
|
|
|
|
let prod = t_vec.iter().product::<u32>();
|
2025-10-26 21:57:53 +08:00
|
|
|
|
Tensor::from_slice(&[prod], 1, t.device())?
|
2025-10-26 21:39:23 +08:00
|
|
|
|
}
|
|
|
|
|
|
DType::I64 => {
|
|
|
|
|
|
let t_vec = t.to_vec1::<i64>()?;
|
|
|
|
|
|
let prod = t_vec.iter().product::<i64>();
|
2025-10-26 21:57:53 +08:00
|
|
|
|
Tensor::from_slice(&[prod], 1, t.device())?
|
2025-10-26 21:39:23 +08:00
|
|
|
|
}
|
|
|
|
|
|
DType::F64 => {
|
|
|
|
|
|
let t_vec = t.to_vec1::<f64>()?;
|
|
|
|
|
|
let prod = t_vec.iter().product::<f64>();
|
2025-10-26 21:57:53 +08:00
|
|
|
|
Tensor::from_slice(&[prod], 1, t.device())?
|
2025-10-26 21:39:23 +08:00
|
|
|
|
}
|
|
|
|
|
|
_ => {
|
|
|
|
|
|
let t_vec = t.to_vec1::<f32>()?;
|
|
|
|
|
|
let prod = t_vec.iter().product::<f32>();
|
2025-10-26 21:57:53 +08:00
|
|
|
|
Tensor::from_slice(&[prod], 1, t.device())?
|
2025-10-26 21:39:23 +08:00
|
|
|
|
}
|
2025-10-26 21:57:53 +08:00
|
|
|
|
}
|
2025-10-26 21:39:23 +08:00
|
|
|
|
}
|
|
|
|
|
|
2 => {
|
|
|
|
|
|
let data_type = t.dtype();
|
2025-10-26 21:57:53 +08:00
|
|
|
|
match data_type {
|
2025-10-26 21:39:23 +08:00
|
|
|
|
DType::U8 => {
|
|
|
|
|
|
let t_vec = t.to_vec2::<u8>()?;
|
|
|
|
|
|
let mut prod_vec = vec![];
|
|
|
|
|
|
for v in t_vec.iter() {
|
|
|
|
|
|
let prod = v.iter().product::<u8>();
|
|
|
|
|
|
prod_vec.push(prod);
|
|
|
|
|
|
}
|
|
|
|
|
|
Tensor::new(prod_vec, t.device())?
|
|
|
|
|
|
}
|
|
|
|
|
|
DType::U32 => {
|
|
|
|
|
|
let t_vec = t.to_vec2::<u32>()?;
|
|
|
|
|
|
let mut prod_vec = vec![];
|
|
|
|
|
|
for v in t_vec.iter() {
|
|
|
|
|
|
let prod = v.iter().product::<u32>();
|
|
|
|
|
|
prod_vec.push(prod);
|
|
|
|
|
|
}
|
|
|
|
|
|
Tensor::new(prod_vec, t.device())?
|
|
|
|
|
|
}
|
|
|
|
|
|
DType::I64 => {
|
|
|
|
|
|
let t_vec = t.to_vec2::<i64>()?;
|
|
|
|
|
|
let mut prod_vec = vec![];
|
|
|
|
|
|
for v in t_vec.iter() {
|
|
|
|
|
|
let prod = v.iter().product::<i64>();
|
|
|
|
|
|
prod_vec.push(prod);
|
|
|
|
|
|
}
|
|
|
|
|
|
Tensor::new(prod_vec, t.device())?
|
|
|
|
|
|
}
|
|
|
|
|
|
DType::F64 => {
|
|
|
|
|
|
let t_vec = t.to_vec2::<f64>()?;
|
|
|
|
|
|
let mut prod_vec = vec![];
|
|
|
|
|
|
for v in t_vec.iter() {
|
|
|
|
|
|
let prod = v.iter().product::<f64>();
|
|
|
|
|
|
prod_vec.push(prod);
|
|
|
|
|
|
}
|
|
|
|
|
|
Tensor::new(prod_vec, t.device())?
|
|
|
|
|
|
}
|
|
|
|
|
|
_ => {
|
|
|
|
|
|
let t_vec = t.to_vec2::<f32>()?;
|
|
|
|
|
|
let mut prod_vec = vec![];
|
|
|
|
|
|
for v in t_vec.iter() {
|
|
|
|
|
|
let prod = v.iter().product::<f32>();
|
|
|
|
|
|
prod_vec.push(prod);
|
|
|
|
|
|
}
|
|
|
|
|
|
Tensor::new(prod_vec, t.device())?
|
|
|
|
|
|
}
|
2025-10-26 21:57:53 +08:00
|
|
|
|
}
|
2025-10-26 21:39:23 +08:00
|
|
|
|
}
|
|
|
|
|
|
_ => {
|
|
|
|
|
|
return Err(anyhow!(format!("can not action this dim")));
|
|
|
|
|
|
}
|
|
|
|
|
|
};
|
|
|
|
|
|
Ok(prod)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2025-10-26 21:57:53 +08:00
|
|
|
|
pub fn mask_index_add(original: &Tensor, mask: &Tensor, add: &Tensor) -> Result<Tensor> {
|
|
|
|
|
|
let visual_nonzero_index = nonzero_index(mask)?;
|
2025-10-26 21:39:23 +08:00
|
|
|
|
let xs = original.index_add(&visual_nonzero_index, add, 0)?;
|
|
|
|
|
|
Ok(xs)
|
|
|
|
|
|
}
|
2025-11-13 00:33:27 +08:00
|
|
|
|
|
|
|
|
|
|
pub fn interpolate_linear(
|
|
|
|
|
|
t: &Tensor,
|
|
|
|
|
|
target_size: usize,
|
|
|
|
|
|
align_corner: Option<bool>,
|
|
|
|
|
|
) -> Result<Tensor> {
|
|
|
|
|
|
// 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::<usize>();
|
|
|
|
|
|
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<Tensor> {
|
|
|
|
|
|
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)
|
|
|
|
|
|
}
|