pub mod cow;
pub mod error;
pub mod layout;
mod strides;
pub mod view;
pub use cow::TensorCow;
pub use error::TensorError;
pub use layout::{ColMajor, Layout, RowMajor};
pub use view::TensorView;
#[cfg(test)]
mod tests {
use super::*;
use strides::row_major_strides;
#[test]
fn test_row_major_strides_3d() {
let s = row_major_strides([2usize, 3, 4]);
assert_eq!(s, [12, 4, 1]);
}
#[test]
fn test_tensor_view_get_2d() {
let data: Vec<i32> = (0..12).collect();
let t = TensorView::<i32, 2>::new(&data, [3, 4]).unwrap();
assert_eq!(t.get([1, 2]).unwrap(), 6);
}
#[test]
fn test_reshape() {
let data: Vec<i32> = (0..12).collect();
let t2d = TensorView::<i32, 2>::new(&data, [3, 4]).unwrap();
let t1d = t2d.reshape([12]).unwrap();
assert_eq!(t1d.num_elements(), 12);
assert_eq!(t1d.get([11]).unwrap(), 11);
}
#[test]
fn test_row_view() {
let data: Vec<f32> = (0..9).map(|x| x as f32).collect();
let t = TensorView::<f32, 2>::new(&data, [3, 3]).unwrap();
let row1 = t.row_view(1).unwrap();
assert_eq!(row1.num_elements(), 3);
assert_eq!(row1.get([0]).unwrap(), 3.0);
assert_eq!(row1.get([2]).unwrap(), 5.0);
}
}