use candela::{Dimension, Tensor};
fn main() {
let a = Tensor::from_slice(&[1.0_f64, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
let b = Tensor::from_slice(&[7.0_f64, 8.0, 9.0, 10.0, 11.0, 12.0], &[3, 2]);
let c = a.matmul(&b).unwrap().materialize();
assert_eq!(c.shape(), &[2, 2]);
assert_eq!(c.data(), &[58.0, 64.0, 139.0, 154.0]);
println!("2D matmul result: {:?}", c.data());
let a = Tensor::from_scalar(1.0_f64, &[2, 3, 4]);
let b = Tensor::from_scalar(1.0_f64, &[2, 4, 5]);
let c = a.matmul(&b).unwrap().materialize();
assert_eq!(c.shape(), &[2, 3, 5]);
assert!(c.data().iter().all(|&x| x == 4.0));
println!("batched matmul shape: {:?}", c.shape());
let a = Tensor::from_scalar(1.0_f64, &[1, 3, 4]);
let b = Tensor::from_scalar(1.0_f64, &[2, 4, 5]);
let c = a.matmul(&b).unwrap().materialize();
assert_eq!(c.shape(), &[2, 3, 5]);
assert!(c.data().iter().all(|&x| x == 4.0));
println!("broadcast lhs batch shape: {:?}", c.shape());
let a = Tensor::from_scalar(1.0_f64, &[2, 3, 4]);
let b = Tensor::from_scalar(1.0_f64, &[1, 4, 5]);
let c = a.matmul(&b).unwrap().materialize();
assert_eq!(c.shape(), &[2, 3, 5]);
assert!(c.data().iter().all(|&x| x == 4.0));
println!("broadcast rhs batch shape: {:?}", c.shape()); }