use super::{Backend, Unsqueezable};
use burn::prelude::{Tensor, Backend as BurnBackend};
use burn::tensor::{BasicOps};
macro_rules! impl_lower_ranked_tensor_ops {
($d:literal) => {
impl <B, K> Unsqueezable for Tensor<B, $d, K>
where B: BurnBackend,
K: BasicOps<B> + 'static {
type Unsqueezed = Tensor<B, {$d + 1}, K>;
fn unsqueeze(&self, dim: usize) -> Self::Unsqueezed {
self.clone().unsqueeze_dim::<{$d + 1}>(dim)
}
}
}
}
macro_rules! impl_core_tensor_ops {
($d:literal) => {
impl <B, K> Backend for Tensor<B, $d, K>
where B: BurnBackend,
K: BasicOps<B> + 'static {
fn shape(&self) -> Vec<usize> {
self.shape().dims.to_vec()
}
fn cat(tensors: &[Self], dim: usize) -> Self {
let owned: Vec<_> = tensors.iter().map(|e| e.clone()).collect();
Tensor::cat(owned, dim)
}
fn eq(&self, other: &Self) -> impl Backend {
let x = self.clone().equal(other.clone());
x
}
fn vectorize_dim(&self, dim: usize) -> Vec<Self> {
let sizes = self.dims();
self.clone().chunk(sizes[dim], dim)
}
fn slice(&self, dimension: usize, seq_start_idx: usize, len: usize) -> Self {
self.clone().narrow(dimension, seq_start_idx, len)
}
fn idx_where_all_true(&self, dim: usize) -> Vec<usize> {
let dims = self.dims();
if dim >= dims.len() {
return vec![]; }
let dim_size = dims[dim];
let mut result = Vec::new();
for i in 0..dim_size {
let slice = self.clone().narrow(dim, i, 1);
let all_tensor = slice.all();
let is_all_true = all_tensor.into_scalar();
if is_all_true {
result.push(i);
}
}
result
}
fn pop(&self, dim: usize, index: usize) -> Self {
let dim_size = self.dims()[dim];
if index >= dim_size {
panic!("Index {} is out of bounds for dimension {} with size {}",
index, dim, dim_size);
}
if index == 0 {
return self.clone().narrow(dim, 1, dim_size - 1);
}
if index == dim_size - 1 {
return self.clone().narrow(dim, 0, dim_size - 1);
}
let first_part = self.clone().narrow(dim, 0, index);
let second_part = self.clone().narrow(dim, index + 1, dim_size - index - 1);
Tensor::cat(vec![first_part, second_part], dim)
}
fn repeat(&self, dim: usize, times: usize) -> Self {
self.clone().repeat_dim(dim, times)
}
}
}
}
impl_lower_ranked_tensor_ops!(1);
impl_lower_ranked_tensor_ops!(2);
impl_lower_ranked_tensor_ops!(3);
impl_lower_ranked_tensor_ops!(4);
impl_lower_ranked_tensor_ops!(5);
impl_lower_ranked_tensor_ops!(6);
impl_lower_ranked_tensor_ops!(7);
impl_lower_ranked_tensor_ops!(8);
impl_core_tensor_ops!(1);
impl_core_tensor_ops!(2);
impl_core_tensor_ops!(3);
impl_core_tensor_ops!(4);
impl_core_tensor_ops!(5);
impl_core_tensor_ops!(6);
impl_core_tensor_ops!(7);
impl_core_tensor_ops!(8);
impl_core_tensor_ops!(9);
#[cfg(test)]
mod test {
use super::{Backend, Unsqueezable};
use burn::backend::ndarray::{NdArray, NdArrayDevice};
use burn::tensor::{Tensor, Float, TensorData};
type BurnBackend = NdArray;
type Device = NdArrayDevice;
type TensorF<const D: usize> = Tensor<BurnBackend, D, Float>;
fn create_sequential_tensor<const D: usize>(shape: &[usize]) -> TensorF<D> {
let size: usize = shape.iter().product();
let data: Vec<f32> = (1..=size).map(|x| x as f32).collect();
let d = burn::tensor::TensorData::new(data, shape);
Tensor::from_data(d, &Device::default())
}
#[test]
fn test_shape() {
let tensor: TensorF<2> = create_sequential_tensor(&[2, 3]);
let shape = Backend::shape(&tensor);
assert_eq!(shape, vec![2, 3]);
}
#[test]
fn test_unsqueeze() {
let tensor: TensorF<2> = create_sequential_tensor(&[2, 3]);
let unsqueezed0 = Unsqueezable::unsqueeze(&tensor, 0);
assert_eq!(Backend::shape(&unsqueezed0), vec![1, 2, 3]);
let unsqueezed1 = Unsqueezable::unsqueeze(&tensor, 1);
assert_eq!(Backend::shape(&unsqueezed1), vec![2, 1, 3]);
let unsqueezed2 = Unsqueezable::unsqueeze(&tensor, 2);
assert_eq!(Backend::shape(&unsqueezed2), vec![2, 3, 1]);
}
#[test]
fn test_cat() {
let tensor1: TensorF<2> = create_sequential_tensor(&[2, 3]);
let tensor2: TensorF<2> = Tensor::from_data(TensorData::new((7..=12).map(|x| x as f32).collect(), &[2, 3]), &Device::default());
let cat_dim0 = Backend::cat(&[tensor1.clone(), tensor2.clone()], 0);
assert_eq!(Backend::shape(&cat_dim0), vec![4, 3]);
let cat_dim1 = Backend::cat(&[tensor1, tensor2], 1);
assert_eq!(Backend::shape(&cat_dim1), vec![2, 6]);
}
#[test]
fn test_vectorize_dim() {
let tensor: TensorF<2> = create_sequential_tensor(&[2, 3]);
let vectorized0 = Backend::vectorize_dim(&tensor, 0);
assert_eq!(vectorized0.len(), 2);
assert_eq!(Backend::shape(&vectorized0[0]), vec![1, 3]);
let tensor2: TensorF<2> = create_sequential_tensor(&[2, 3]);
let vectorized1 = Backend::vectorize_dim(&tensor2, 1);
assert_eq!(vectorized1.len(), 3);
assert_eq!(Backend::shape(&vectorized1[0]), vec![2, 1]);
}
#[test]
fn test_slice() {
let tensor: TensorF<3> = create_sequential_tensor(&[2, 3, 2]);
let slice0 = Backend::slice(&tensor, 0, 0, 1);
assert_eq!(Backend::shape(&slice0), vec![1, 3, 2]);
let slice1 = Backend::slice(&tensor, 1, 1, 2);
assert_eq!(Backend::shape(&slice1), vec![2, 2, 2]);
let slice2 = Backend::slice(&tensor, 2, 0, 1);
assert_eq!(Backend::shape(&slice2), vec![2, 3, 1]);
}
#[test]
fn test_idx_where_all_true() {
let bool_tensor1: TensorF<2> = Tensor::from_data(
[ [1.0, 1.0], [0.0, 1.0], [1.0, 0.0]], &Device::default());
eprintln!("{:?}", bool_tensor1.shape());
let indices0 = Backend::idx_where_all_true(&bool_tensor1, 0);
assert_eq!(indices0, vec![0]);
let indices1 = Backend::idx_where_all_true(&bool_tensor1, 1);
assert_eq!(indices1.len(), 0);
let bool_tensor2: TensorF<2> = Tensor::from_data(
[ [1.0, 0.0], [1.0, 0.0], [1.0, 1.0]], &Device::default());
let indices2 = Backend::idx_where_all_true(&bool_tensor2, 0);
assert_eq!(indices2, vec![2]);
let indices3 = Backend::idx_where_all_true(&bool_tensor2, 1);
assert_eq!(indices3, vec![0]);
}
#[test]
fn test_pop() {
let tensor: TensorF<2> = create_sequential_tensor(&[2, 3]);
let popped0 = Backend::pop(&tensor, 0, 0);
assert_eq!(Backend::shape(&popped0), vec![1, 3]);
let tensor: TensorF<2> = create_sequential_tensor(&[2, 3]);
let popped1 = Backend::pop(&tensor, 1, 1);
assert_eq!(Backend::shape(&popped1), vec![2, 2]);
let tensor: TensorF<2> = create_sequential_tensor(&[2, 3]);
let popped2 = Backend::pop(&tensor, 1, 2);
assert_eq!(Backend::shape(&popped2), vec![2, 2]);
}
#[test]
fn test_repeat() {
let tensor: TensorF<2> = Tensor::from_data([[1.0, 2.0, 3.0]], &Device::default());
let repeated0 = Backend::repeat(&tensor, 0, 2);
assert_eq!(Backend::shape(&repeated0), vec![2, 3]);
let tensor2: TensorF<2> = Tensor::from_data([[1.0, 2.0], [3.0, 4.0]], &Device::default());
let repeated1 = Backend::repeat(&tensor2, 1, 3);
assert_eq!(Backend::shape(&repeated1), vec![2, 6]);
}
#[test]
#[should_panic(expected = "Index 3 is out of bounds for dimension 0 with size 2")]
fn test_pop_out_of_bounds() {
let tensor: TensorF<2> = create_sequential_tensor(&[2, 3]);
let _ = Backend::pop(&tensor, 0, 3);
}
#[test]
fn test_tensor_3d() {
let tensor: TensorF<3> = create_sequential_tensor(&[2, 3, 2]);
assert_eq!(Backend::shape(&tensor), vec![2, 3, 2]);
let unsqueezed = Unsqueezable::unsqueeze(&tensor, 1);
assert_eq!(Backend::shape(&unsqueezed), vec![2, 1, 3, 2]);
let sliced = Backend::slice(&tensor, 0, 1, 1);
assert_eq!(Backend::shape(&sliced), vec![1, 3, 2]);
let vectorized = Backend::vectorize_dim(&tensor, 0);
assert_eq!(vectorized.len(), 2);
assert_eq!(Backend::shape(&vectorized[0]), vec![1, 3, 2]);
}
}