mod common;
use candela::{OpError, Tensor, s};
#[test]
fn get_1d() {
let t = Tensor::from_slice(&[10.0, 20.0, 30.0], &[3]);
assert_eq!(t.get(&[0]).unwrap(), &10.0);
assert_eq!(t.get(&[2]).unwrap(), &30.0);
}
#[test]
fn get_2d() {
let t = Tensor::from_slice(&[0.0, 1.0, 2.0, 3.0, 4.0, 5.0], &[2, 3]);
assert_eq!(t.get(&[0, 0]).unwrap(), &0.0);
assert_eq!(t.get(&[1, 2]).unwrap(), &5.0);
}
#[test]
fn get_wrong_rank() {
let t = Tensor::from_slice(&[1.0, 2.0, 3.0], &[3]);
assert!(matches!(t.get(&[0, 0]), Err(OpError::NotEnoughAxes(1, 2))));
}
#[test]
fn get_out_of_bounds() {
let t = Tensor::from_slice(&[1.0, 2.0, 3.0], &[3]);
assert!(matches!(t.get(&[3]), Err(OpError::IndexOutOfBounds)));
}
#[test]
fn get_sliced() {
let t = Tensor::from_slice(&[0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0], &[3, 3]);
let sliced = t.slice(s![1..3, ..]).unwrap().materialize();
assert_eq!(sliced.get(&[0, 1]).unwrap(), &4.0);
assert_eq!(sliced.get(&[1, 2]).unwrap(), &8.0);
}
#[test]
fn get_transposed() {
let t = Tensor::from_slice(&[0.0, 1.0, 2.0, 3.0, 4.0, 5.0], &[2, 3]);
let tr = t.transpose().materialize();
assert_eq!(tr.get(&[2, 1]).unwrap(), &5.0);
}
#[test]
fn index_1d() {
let t = Tensor::from_slice(&[7.0, 8.0, 9.0], &[3]);
assert_eq!(t[&[1][..]], 8.0);
}
#[test]
fn index_2d() {
let t = Tensor::from_slice(&[1.0, 2.0, 3.0, 4.0], &[2, 2]);
assert_eq!(t[&[1, 0][..]], 3.0);
}
#[test]
fn item_scalar() {
let t = Tensor::from_scalar(42.0_f64, &[1]);
assert_eq!(t.item(), &42.0);
}
#[test]
fn item_after_materialize() {
let t = Tensor::from_scalar(3.0_f64, &[1]);
let result = (t * 7.0).materialize();
assert_eq!(result.item(), &21.0);
}