mod common;
use candela::{Dimension, Tensor, arange, s};
use common::assert_approx_eq;
#[test]
fn view_is_zero_copy() {
let t: Tensor<f64> = arange!(12);
let viewed = t.view(&[3, 4]).unwrap().materialize();
assert_eq!(viewed.data().as_ptr(), t.data().as_ptr());
}
#[test]
fn transpose_is_zero_copy() {
let t: Tensor<f64> = arange!(12);
let viewed = t.view(&[3, 4]).unwrap();
let t2 = viewed.materialize();
let transposed = t2.transpose().materialize();
assert_eq!(transposed.data().as_ptr(), t2.data().as_ptr());
}
#[test]
fn slice_is_zero_copy() {
let t: Tensor<f64> = arange!(12);
let viewed = t.view(&[3, 4]).unwrap().materialize();
let sliced = viewed.slice(s![0..2, ..]).unwrap().materialize();
assert_eq!(sliced.data().as_ptr(), viewed.data().as_ptr());
}
#[test]
fn transpose_2x2() {
let t = Tensor::from_slice(&[1.0, 2.0, 3.0, 4.0], &[2, 2]);
let result = (t.transpose() * 1.0).materialize();
assert_eq!(result.shape(), &[2, 2]);
assert_approx_eq(result.data(), &[1.0, 3.0, 2.0, 4.0]);
}
#[test]
fn transpose_2x3() {
let t = Tensor::from_slice(&[0.0, 1.0, 2.0, 3.0, 4.0, 5.0], &[2, 3]);
let result = (t.transpose() * 1.0).materialize();
assert_eq!(result.shape(), &[3, 2]);
assert_approx_eq(result.data(), &[0.0, 3.0, 1.0, 4.0, 2.0, 5.0]);
}
#[test]
fn slice_then_add() {
let t = arange!(5);
let sliced = t.slice(s![2..4]).unwrap();
let result = (sliced + 1.0).materialize();
assert_approx_eq(result.data(), &[3.0, 4.0]);
}
#[test]
fn slice_2d_row_range() {
let t = Tensor::from_slice(
&[0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0],
&[3, 4],
);
let sliced = t.slice(s![1..3, ..]).unwrap().materialize();
let temp: Box<[f64]> = sliced.iter().cloned().collect();
assert_eq!(sliced.shape(), &[2, 4]);
assert_approx_eq(&temp, &[4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0]);
}
#[test]
fn as_contiguous_transposed() {
let t = Tensor::from_slice(&[0.0, 1.0, 2.0, 3.0], &[2, 2]);
let transposed = t.transpose();
let result = transposed.as_contiguous().materialize();
assert!(result.is_contiguous());
assert_approx_eq(result.data(), &[0.0, 2.0, 1.0, 3.0]);
}
#[test]
fn view_after_as_contiguous() {
let t = Tensor::from_slice(&[0.0, 1.0, 2.0, 3.0, 4.0, 5.0], &[2, 3]);
let cont = t.transpose().as_contiguous().materialize();
let viewed = cont.view(&[6]).unwrap().materialize();
assert_eq!(viewed.shape(), &[6]);
}