use matten::Tensor;
fn render_block(label: &str, t: &Tensor) -> String {
if t.shape().len() <= 2 {
format!("{label:<16} shape={:?}\n{t}", t.shape())
} else {
format!("{label:<16} {t}")
}
}
fn main() {
let a = Tensor::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
let b = Tensor::new(vec![10.0, 20.0, 30.0], &[3]);
println!("== Broadcasting ==");
let input_a = render_block("input A", &a);
assert_eq!(
input_a,
"input A shape=[2, 3]\n1.0 2.0 3.0\n4.0 5.0 6.0"
);
println!("{input_a}");
let input_b = render_block("input b", &b);
assert_eq!(input_b, "input b shape=[3]\n10.0 20.0 30.0");
println!("{input_b}");
let broadcast = &a + &b;
let broadcast_block = render_block("A + b", &broadcast);
assert_eq!(
broadcast_block,
"A + b shape=[2, 3]\n11.0 22.0 33.0\n14.0 25.0 36.0"
);
println!("{broadcast_block}");
println!("meaning b repeats across rows");
assert_eq!(broadcast.shape(), &[2, 3]);
assert_eq!(broadcast.as_slice(), &[11.0, 22.0, 33.0, 14.0, 25.0, 36.0]);
println!();
println!("== Reshape ==");
let reshaped = a.reshape(&[3, 2]);
let reshape_input = render_block("[2, 3] input", &a);
assert_eq!(
reshape_input,
"[2, 3] input shape=[2, 3]\n1.0 2.0 3.0\n4.0 5.0 6.0"
);
println!("{reshape_input}");
let reshape_view = render_block("[3, 2] view", &reshaped);
assert_eq!(
reshape_view,
"[3, 2] view shape=[3, 2]\n1.0 2.0\n3.0 4.0\n5.0 6.0"
);
println!("{reshape_view}");
println!("meaning row-major values stay in the same order");
assert_eq!(reshaped.shape(), &[3, 2]);
assert_eq!(reshaped.as_slice(), a.as_slice());
println!();
println!("== Axis reductions ==");
let col_means = a.mean_axis(0);
let row_means = a.mean_axis(1);
println!(
"mean_axis(0) collapse rows, keep columns -> shape {:?}, values {:?}",
col_means.shape(),
col_means.as_slice()
);
println!(
"mean_axis(1) collapse columns, keep rows -> shape {:?}, values {:?}",
row_means.shape(),
row_means.as_slice()
);
assert_eq!(col_means.shape(), &[3]);
assert_eq!(col_means.as_slice(), &[2.5, 3.5, 4.5]);
assert_eq!(row_means.shape(), &[2]);
assert_eq!(row_means.as_slice(), &[2.0, 5.0]);
println!();
println!("== Matrix multiplication ==");
let left = Tensor::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
let right = Tensor::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[3, 2]);
let product = left.matmul(&right);
let left_block = render_block("left", &left);
assert_eq!(
left_block,
"left shape=[2, 3]\n1.0 2.0 3.0\n4.0 5.0 6.0"
);
println!("{left_block}");
let right_block = render_block("right", &right);
assert_eq!(
right_block,
"right shape=[3, 2]\n1.0 2.0\n3.0 4.0\n5.0 6.0"
);
println!("{right_block}");
let product_block = render_block("left.matmul", &product);
assert_eq!(
product_block,
"left.matmul shape=[2, 2]\n22.0 28.0\n49.0 64.0"
);
println!("{product_block}");
println!("meaning [2, 3] x [3, 2] -> [2, 2]");
assert_eq!(product.shape(), &[2, 2]);
assert_eq!(product.as_slice(), &[22.0, 28.0, 49.0, 64.0]);
println!();
println!("57_visual_shape_axis_summary: OK");
}