use crate::ADD;
use crate::Tensor4;
use russell_lab::{small_mat_add, small_mat_mat_mul};
pub fn t4_add<const N: usize>(c: &mut Tensor4<N>, alpha: f64, a: &Tensor4<N>, beta: f64, b: &Tensor4<N>) {
small_mat_add(&mut c.mat, alpha, &a.mat, beta, &b.mat, N);
}
pub fn t4_ddot_t4<const N: usize>(ee: &mut Tensor4<N>, op: u8, alpha: f64, cc: &Tensor4<N>, dd: &Tensor4<N>) {
let beta = if op == ADD { 1.0 } else { 0.0 };
small_mat_mat_mul(&mut ee.mat, alpha, &cc.mat, &dd.mat, beta, N);
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{ADD, SET, SamplesTensor4};
use russell_lab::{Matrix, mat_approx_eq};
#[test]
fn t4_add_works() {
let mut a = Tensor4::<4>::new();
let mut b = Tensor4::<4>::new();
let mut c = Tensor4::<4>::new();
a.sym_set_std(0, 0, 0, 0, 1.0);
b.sym_set_std(0, 0, 0, 0, 1.0);
t4_add(&mut c, 2.0, &a, 3.0, &b);
#[rustfmt::skip]
let correct = &[
[5.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
];
mat_approx_eq(&c.as_std_matrix(), correct, 1e-14);
}
#[test]
fn t4_ddot_t4_set_works() {
let cc = Tensor4::<4>::from_std_matrix(&SamplesTensor4::SYM_2D_SAMPLE1_STD_MATRIX).unwrap();
let mut ee = Tensor4::<4>::new();
t4_ddot_t4(&mut ee, SET, 2.0, &cc, &cc);
let out = ee.as_std_matrix();
assert_eq!(
format!("{:.1}", out),
"┌ ┐\n\
│ 820.0 872.0 924.0 1288.0 0.0 0.0 1288.0 0.0 0.0 │\n\
│ 1120.0 1202.0 1284.0 1858.0 0.0 0.0 1858.0 0.0 0.0 │\n\
│ 1420.0 1532.0 1644.0 2428.0 0.0 0.0 2428.0 0.0 0.0 │\n\
│ 2620.0 2852.0 3084.0 4708.0 0.0 0.0 4708.0 0.0 0.0 │\n\
│ 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 │\n\
│ 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 │\n\
│ 2620.0 2852.0 3084.0 4708.0 0.0 0.0 4708.0 0.0 0.0 │\n\
│ 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 │\n\
│ 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 │\n\
└ ┘"
);
}
#[test]
fn t4_ddot_t4_add_works() {
let cc = Tensor4::<4>::from_std_matrix(&SamplesTensor4::SYM_2D_SAMPLE1_STD_MATRIX).unwrap();
let mut mat = Matrix::new(9, 9);
mat.set(0, 0, 0.1);
mat.set(1, 1, 0.1);
mat.set(2, 2, 0.1);
let mut ee = Tensor4::<4>::from_std_matrix(&mat).unwrap();
t4_ddot_t4(&mut ee, ADD, 2.0, &cc, &cc);
let out = ee.as_std_matrix();
assert_eq!(
format!("{:.1}", out),
"┌ ┐\n\
│ 820.1 872.0 924.0 1288.0 0.0 0.0 1288.0 0.0 0.0 │\n\
│ 1120.0 1202.1 1284.0 1858.0 0.0 0.0 1858.0 0.0 0.0 │\n\
│ 1420.0 1532.0 1644.1 2428.0 0.0 0.0 2428.0 0.0 0.0 │\n\
│ 2620.0 2852.0 3084.0 4708.0 0.0 0.0 4708.0 0.0 0.0 │\n\
│ 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 │\n\
│ 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 │\n\
│ 2620.0 2852.0 3084.0 4708.0 0.0 0.0 4708.0 0.0 0.0 │\n\
│ 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 │\n\
│ 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 │\n\
└ ┘"
);
}
}