use crate::ADD;
use crate::{Tensor2, Tensor4};
pub fn t2_dyad_t2<const N: usize>(dd: &mut Tensor4<N>, op: u8, alpha: f64, a: &Tensor2<N>, b: &Tensor2<N>) {
if op == ADD {
for m in 0..N {
for n in 0..N {
dd.add(m, n, alpha * a.vec[m] * b.vec[n]);
}
}
} else {
for m in 0..N {
for n in 0..N {
dd.set(m, n, alpha * a.vec[m] * b.vec[n]);
}
}
}
}
pub fn t4_ddot_t2<const N: usize>(b: &mut Tensor2<N>, op: u8, alpha: f64, dd: &Tensor4<N>, a: &Tensor2<N>) {
if op == ADD {
for m in 0..N {
let mut s = 0.0;
for n in 0..N {
s += dd.get(m, n) * a.vec[n];
}
b.vec[m] += alpha * s;
}
} else {
for m in 0..N {
let mut s = 0.0;
for n in 0..N {
s += dd.get(m, n) * a.vec[n];
}
b.vec[m] = alpha * s;
}
}
}
pub fn t2_ddot_t4<const N: usize>(b: &mut Tensor2<N>, op: u8, alpha: f64, a: &Tensor2<N>, dd: &Tensor4<N>) {
if op == ADD {
for n in 0..N {
let mut s = 0.0;
for m in 0..N {
s += a.vec[m] * dd.get(m, n);
}
b.vec[n] += alpha * s;
}
} else {
for n in 0..N {
let mut s = 0.0;
for m in 0..N {
s += a.vec[m] * dd.get(m, n);
}
b.vec[n] = alpha * s;
}
}
}
pub fn t2_ddot_t4_ddot_t2<const N: usize>(a: &Tensor2<N>, dd: &Tensor4<N>, b: &Tensor2<N>) -> f64 {
let mut s = 0.0;
for m in 0..N {
for n in 0..N {
s += a.vec[m] * dd.get(m, n) * b.vec[n];
}
}
s
}
pub fn t4_ddot_t2_dyad_t2_ddot_t4<const N: usize>(
ee: &mut Tensor4<N>,
alpha: f64,
dd: &Tensor4<N>,
beta: f64,
a: &Tensor2<N>,
b: &Tensor2<N>,
) {
for m in 0..N {
for n in 0..N {
ee.set(m, n, alpha * dd.get(m, n));
for p in 0..N {
for q in 0..N {
ee.set(
m,
n,
ee.get(m, n) + beta * dd.get(m, p) * a.vec[p] * b.vec[q] * dd.get(q, n),
);
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{MN_TO_IJKL, SET, SamplesTensor4};
use russell_lab::{Matrix, approx_eq, mat_approx_eq};
#[test]
fn t2_dyad_t2_works() {
#[rustfmt::skip]
let a = Tensor2::<9>::from_std_matrix(&[
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
[7.0, 8.0, 9.0],
]).unwrap();
#[rustfmt::skip]
let b = Tensor2::<9>::from_std_matrix(&[
[0.5, 0.5, 0.5],
[0.5, 0.5, 0.5],
[0.5, 0.5, 0.5],
]).unwrap();
let mut dd = Tensor4::<9>::new();
t2_dyad_t2(&mut dd, SET, 2.0, &a, &b);
let mat = dd.as_std_matrix();
assert_eq!(
format!("{:.1}", mat),
"┌ ┐\n\
│ 1.0 1.0 1.0 1.0 1.0 1.0 1.0 1.0 1.0 │\n\
│ 5.0 5.0 5.0 5.0 5.0 5.0 5.0 5.0 5.0 │\n\
│ 9.0 9.0 9.0 9.0 9.0 9.0 9.0 9.0 9.0 │\n\
│ 2.0 2.0 2.0 2.0 2.0 2.0 2.0 2.0 2.0 │\n\
│ 6.0 6.0 6.0 6.0 6.0 6.0 6.0 6.0 6.0 │\n\
│ 3.0 3.0 3.0 3.0 3.0 3.0 3.0 3.0 3.0 │\n\
│ 4.0 4.0 4.0 4.0 4.0 4.0 4.0 4.0 4.0 │\n\
│ 8.0 8.0 8.0 8.0 8.0 8.0 8.0 8.0 8.0 │\n\
│ 7.0 7.0 7.0 7.0 7.0 7.0 7.0 7.0 7.0 │\n\
└ ┘"
);
#[rustfmt::skip]
let a = Tensor2::<6>::from_std_matrix(&[
[1.0, 2.0, 3.0],
[2.0, 5.0, 6.0],
[3.0, 6.0, 9.0],
]).unwrap();
#[rustfmt::skip]
let b = Tensor2::<6>::from_std_matrix(&[
[0.5, 0.5, 0.5],
[0.5, 0.5, 0.5],
[0.5, 0.5, 0.5],
]).unwrap();
let mut dd = Tensor4::<6>::new();
t2_dyad_t2(&mut dd, SET, 2.0, &a, &b);
let mat = dd.as_std_matrix();
assert_eq!(
format!("{:.1}", mat),
"┌ ┐\n\
│ 1.0 1.0 1.0 1.0 1.0 1.0 1.0 1.0 1.0 │\n\
│ 5.0 5.0 5.0 5.0 5.0 5.0 5.0 5.0 5.0 │\n\
│ 9.0 9.0 9.0 9.0 9.0 9.0 9.0 9.0 9.0 │\n\
│ 2.0 2.0 2.0 2.0 2.0 2.0 2.0 2.0 2.0 │\n\
│ 6.0 6.0 6.0 6.0 6.0 6.0 6.0 6.0 6.0 │\n\
│ 3.0 3.0 3.0 3.0 3.0 3.0 3.0 3.0 3.0 │\n\
│ 2.0 2.0 2.0 2.0 2.0 2.0 2.0 2.0 2.0 │\n\
│ 6.0 6.0 6.0 6.0 6.0 6.0 6.0 6.0 6.0 │\n\
│ 3.0 3.0 3.0 3.0 3.0 3.0 3.0 3.0 3.0 │\n\
└ ┘"
);
#[rustfmt::skip]
let a = Tensor2::<4>::from_std_matrix(&[
[1.0, 2.0, 0.0],
[2.0, 5.0, 0.0],
[0.0, 0.0, 9.0],
]).unwrap();
#[rustfmt::skip]
let b = Tensor2::<4>::from_std_matrix(&[
[0.5, 0.5, 0.0],
[0.5, 0.5, 0.0],
[0.0, 0.0, 0.5],
]).unwrap();
let mut dd = Tensor4::<4>::new();
t2_dyad_t2(&mut dd, SET, 2.0, &a, &b);
let mat = dd.as_std_matrix();
assert_eq!(
format!("{:.1}", mat),
"┌ ┐\n\
│ 1.0 1.0 1.0 1.0 0.0 0.0 1.0 0.0 0.0 │\n\
│ 5.0 5.0 5.0 5.0 0.0 0.0 5.0 0.0 0.0 │\n\
│ 9.0 9.0 9.0 9.0 0.0 0.0 9.0 0.0 0.0 │\n\
│ 2.0 2.0 2.0 2.0 0.0 0.0 2.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\
│ 2.0 2.0 2.0 2.0 0.0 0.0 2.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 t2_dyad_t2_add_works() {
#[rustfmt::skip]
let a = Tensor2::<9>::from_std_matrix(&[
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
[7.0, 8.0, 9.0],
]).unwrap();
#[rustfmt::skip]
let b = Tensor2::<9>::from_std_matrix(&[
[0.5, 0.5, 0.5],
[0.5, 0.5, 0.5],
[0.5, 0.5, 0.5],
]).unwrap();
let mat = Matrix::filled(9, 9, 0.1);
let mut dd = Tensor4::<9>::from_std_matrix(&mat).unwrap();
t2_dyad_t2(&mut dd, ADD, 2.0, &a, &b);
let mat = dd.as_std_matrix();
let correct = "┌ ┐\n\
│ 1.1 1.1 1.1 1.1 1.1 1.1 1.1 1.1 1.1 │\n\
│ 5.1 5.1 5.1 5.1 5.1 5.1 5.1 5.1 5.1 │\n\
│ 9.1 9.1 9.1 9.1 9.1 9.1 9.1 9.1 9.1 │\n\
│ 2.1 2.1 2.1 2.1 2.1 2.1 2.1 2.1 2.1 │\n\
│ 6.1 6.1 6.1 6.1 6.1 6.1 6.1 6.1 6.1 │\n\
│ 3.1 3.1 3.1 3.1 3.1 3.1 3.1 3.1 3.1 │\n\
│ 4.1 4.1 4.1 4.1 4.1 4.1 4.1 4.1 4.1 │\n\
│ 8.1 8.1 8.1 8.1 8.1 8.1 8.1 8.1 8.1 │\n\
│ 7.1 7.1 7.1 7.1 7.1 7.1 7.1 7.1 7.1 │\n\
└ ┘";
assert_eq!(format!("{:.1}", mat), correct);
}
fn check_dyad<const N: usize>(s: f64, a_ten: &Tensor2<N>, b_ten: &Tensor2<N>, dd_ten: &Tensor4<N>, tol: f64) {
let a = a_ten.as_std_matrix();
let b = b_ten.as_std_matrix();
let dd = dd_ten.as_std_matrix();
let mut correct = Matrix::new(9, 9); for m in 0..9 {
for n in 0..9 {
let (i, j, k, l) = MN_TO_IJKL[m][n];
correct.set(m, n, s * a.get(i, j) * b.get(k, l));
}
}
mat_approx_eq(&dd, &correct, tol);
}
#[test]
fn t2_dyad_t2_works_extra() {
#[rustfmt::skip]
let a = Tensor2::<9>::from_std_matrix(&[
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
[7.0, 8.0, 9.0],
]).unwrap();
#[rustfmt::skip]
let b = Tensor2::<9>::from_std_matrix(&[
[9.0, 8.0, 7.0],
[6.0, 5.0, 4.0],
[3.0, 2.0, 1.0],
]).unwrap();
let mut dd = Tensor4::<9>::new();
t2_dyad_t2(&mut dd, SET, 2.0, &a, &b);
let mat = dd.as_std_matrix();
let correct = Matrix::from(&[
[18.0, 10.0, 2.0, 16.0, 8.0, 14.0, 12.0, 4.0, 6.0],
[90.0, 50.0, 10.0, 80.0, 40.0, 70.0, 60.0, 20.0, 30.0],
[162.0, 90.0, 18.0, 144.0, 72.0, 126.0, 108.0, 36.0, 54.0],
[36.0, 20.0, 4.0, 32.0, 16.0, 28.0, 24.0, 8.0, 12.0],
[108.0, 60.0, 12.0, 96.0, 48.0, 84.0, 72.0, 24.0, 36.0],
[54.0, 30.0, 6.0, 48.0, 24.0, 42.0, 36.0, 12.0, 18.0],
[72.0, 40.0, 8.0, 64.0, 32.0, 56.0, 48.0, 16.0, 24.0],
[144.0, 80.0, 16.0, 128.0, 64.0, 112.0, 96.0, 32.0, 48.0],
[126.0, 70.0, 14.0, 112.0, 56.0, 98.0, 84.0, 28.0, 42.0],
]);
mat_approx_eq(&mat, &correct, 1e-13);
check_dyad(2.0, &a, &b, &dd, 1e-13);
#[rustfmt::skip]
let a = Tensor2::<6>::from_std_matrix(&[
[1.0, 4.0, 6.0],
[4.0, 2.0, 5.0],
[6.0, 5.0, 3.0],
]).unwrap();
#[rustfmt::skip]
let b = Tensor2::<6>::from_std_matrix(&[
[3.0, 5.0, 6.0],
[5.0, 2.0, 4.0],
[6.0, 4.0, 1.0],
]).unwrap();
let mut dd = Tensor4::<6>::new();
t2_dyad_t2(&mut dd, SET, 2.0, &a, &b);
let mat = dd.as_std_matrix();
let correct = Matrix::from(&[
[6.0, 4.0, 2.0, 10.0, 8.0, 12.0, 10.0, 8.0, 12.0],
[12.0, 8.0, 4.0, 20.0, 16.0, 24.0, 20.0, 16.0, 24.0],
[18.0, 12.0, 6.0, 30.0, 24.0, 36.0, 30.0, 24.0, 36.0],
[24.0, 16.0, 8.0, 40.0, 32.0, 48.0, 40.0, 32.0, 48.0],
[30.0, 20.0, 10.0, 50.0, 40.0, 60.0, 50.0, 40.0, 60.0],
[36.0, 24.0, 12.0, 60.0, 48.0, 72.0, 60.0, 48.0, 72.0],
[24.0, 16.0, 8.0, 40.0, 32.0, 48.0, 40.0, 32.0, 48.0],
[30.0, 20.0, 10.0, 50.0, 40.0, 60.0, 50.0, 40.0, 60.0],
[36.0, 24.0, 12.0, 60.0, 48.0, 72.0, 60.0, 48.0, 72.0],
]);
mat_approx_eq(&mat, &correct, 1e-13);
check_dyad(2.0, &a, &b, &dd, 1e-13);
#[rustfmt::skip]
let a = Tensor2::<4>::from_std_matrix(&[
[1.0, 4.0, 0.0],
[4.0, 2.0, 0.0],
[0.0, 0.0, 3.0],
]).unwrap();
#[rustfmt::skip]
let b = Tensor2::<4>::from_std_matrix(&[
[3.0, 4.0, 0.0],
[4.0, 2.0, 0.0],
[0.0, 0.0, 1.0],
]).unwrap();
let mut dd = Tensor4::<4>::new();
t2_dyad_t2(&mut dd, SET, 2.0, &a, &b);
let mat = dd.as_std_matrix();
let correct = Matrix::from(&[
[6.0, 4.0, 2.0, 8.0, 0.0, 0.0, 8.0, 0.0, 0.0],
[12.0, 8.0, 4.0, 16.0, 0.0, 0.0, 16.0, 0.0, 0.0],
[18.0, 12.0, 6.0, 24.0, 0.0, 0.0, 24.0, 0.0, 0.0],
[24.0, 16.0, 8.0, 32.0, 0.0, 0.0, 32.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],
[24.0, 16.0, 8.0, 32.0, 0.0, 0.0, 32.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(&mat, &correct, 1e-14);
check_dyad(2.0, &a, &b, &dd, 1e-15);
}
#[test]
fn t4_ddot_t2_works() {
let dd = Tensor4::<4>::from_std_matrix(&SamplesTensor4::SYM_2D_SAMPLE1_STD_MATRIX).unwrap();
#[rustfmt::skip]
let a = Tensor2::<4>::from_std_matrix(&[
[-1.0, -2.0, 0.0],
[-2.0, 2.0, 0.0],
[ 0.0, 0.0, -3.0]]).unwrap();
let mut b = Tensor2::<4>::new();
t4_ddot_t2(&mut b, SET, 1.0, &dd, &a);
let out = b.as_std_matrix();
assert_eq!(
format!("{:.1}", out),
"┌ ┐\n\
│ -46.0 -154.0 0.0 │\n\
│ -154.0 -64.0 0.0 │\n\
│ 0.0 0.0 -82.0 │\n\
└ ┘"
);
}
#[test]
fn t4_ddot_t2_add_works() {
let dd = Tensor4::<4>::from_std_matrix(&SamplesTensor4::SYM_2D_SAMPLE1_STD_MATRIX).unwrap();
#[rustfmt::skip]
let a = Tensor2::<4>::from_std_matrix(&[
[-1.0, -2.0, 0.0],
[-2.0, 2.0, 0.0],
[ 0.0, 0.0, -3.0],
]).unwrap();
#[rustfmt::skip]
let mut b = Tensor2::<4>::from_std_matrix(&[
[-2000.0, -2000.0, 0.0],
[-2000.0, -2000.0, 0.0],
[ 0.0, 0.0, -2000.0],
]).unwrap();
t4_ddot_t2(&mut b, ADD, 1.0, &dd, &a);
let out = b.as_std_matrix();
assert_eq!(
format!("{:.1}", out),
"┌ ┐\n\
│ -2046.0 -2154.0 0.0 │\n\
│ -2154.0 -2064.0 0.0 │\n\
│ 0.0 0.0 -2082.0 │\n\
└ ┘"
);
}
#[test]
fn t2_ddot_t4_works() {
let dd = Tensor4::<4>::from_std_matrix(&SamplesTensor4::SYM_2D_SAMPLE1_STD_MATRIX).unwrap();
#[rustfmt::skip]
let a = Tensor2::<4>::from_std_matrix(&[
[-1.0, -2.0, 0.0],
[-2.0, 2.0, 0.0],
[ 0.0, 0.0, -3.0]]).unwrap();
let mut b = Tensor2::<4>::new();
t2_ddot_t4(&mut b, SET, 1.0, &a, &dd);
let out = b.as_std_matrix();
assert_eq!(
format!("{:.1}", out),
"┌ ┐\n\
│ -90.0 -144.0 0.0 │\n\
│ -144.0 -96.0 0.0 │\n\
│ 0.0 0.0 -102.0 │\n\
└ ┘"
);
}
#[test]
fn t2_ddot_t4_add_works() {
let dd = Tensor4::<4>::from_std_matrix(&SamplesTensor4::SYM_2D_SAMPLE1_STD_MATRIX).unwrap();
#[rustfmt::skip]
let a = Tensor2::<4>::from_std_matrix(&[
[-1.0, -2.0, 0.0],
[-2.0, 2.0, 0.0],
[ 0.0, 0.0, -3.0]]).unwrap();
let mut b =
Tensor2::<4>::from_std_matrix(&[[1000.0, 0.0, 0.0], [0.0, 2000.0, 0.0], [0.0, 0.0, 3000.0]]).unwrap();
t2_ddot_t4(&mut b, ADD, 1.0, &a, &dd);
let correct = &[[910.0, -144.0, 0.0], [-144.0, 1904.0, 0.0], [0.0, 0.0, 2898.0]];
mat_approx_eq(&b.as_std_matrix(), correct, 1e-13);
}
#[test]
fn t2_ddot_t4_ddot_t2_works() {
#[rustfmt::skip]
let a = Tensor2::<9>::from_std_matrix(&[
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
[7.0, 8.0, 9.0],
]).unwrap();
#[rustfmt::skip]
let b = Tensor2::<9>::from_std_matrix(&[
[9.0, 8.0, 7.0],
[6.0, 5.0, 4.0],
[3.0, 2.0, 1.0],
]).unwrap();
let mat = Matrix::filled(9, 9, -1.0);
let dd = Tensor4::<9>::from_std_matrix(&mat).unwrap();
let s = t2_ddot_t4_ddot_t2(&a, &dd, &b);
approx_eq(s, -2025.0, 1e-15);
}
#[test]
fn t4_ddot_t2_dyad_t2_ddot_t4_works1() {
#[rustfmt::skip]
let a = Tensor2::<9>::from_std_matrix(&[
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
[7.0, 8.0, 9.0],
]).unwrap();
#[rustfmt::skip]
let b = Tensor2::<9>::from_std_matrix(&[
[9.0, 8.0, 7.0],
[6.0, 5.0, 4.0],
[3.0, 2.0, 1.0],
]).unwrap();
let mat = Matrix::filled(9, 9, -1.0);
let dd = Tensor4::<9>::from_std_matrix(&mat).unwrap();
let mut ee = Tensor4::<9>::new();
t4_ddot_t2_dyad_t2_ddot_t4(&mut ee, 2.0, &dd, 3.0, &a, &b);
let correct = [
[6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073.],
[6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073.],
[6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073.],
[6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073.],
[6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073.],
[6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073.],
[6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073.],
[6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073.],
[6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073., 6073.],
];
mat_approx_eq(&ee.as_std_matrix(), &correct, 1e-15);
}
}