use mdarray::{Array, Dim, Layout, Shape, Slice, tensor};
use num_complex::ComplexFloat;
use num_traits::{One, Zero};
pub fn pretty_print<T: ComplexFloat + std::fmt::Display, D0: Dim, D1: Dim>(mat: &Array<T, (D0, D1)>)
where
<T as num_complex::ComplexFloat>::Real: std::fmt::Display,
{
let shape = mat.shape();
for i in 0..shape.dim(0) {
for j in 0..shape.dim(1) {
let v = mat[[i, j]];
print!("{:>10.4} {:+.4}i ", v.re(), v.im(),);
}
println!();
}
println!();
}
#[doc(hidden)]
pub fn into_i32<T>(x: T) -> i32
where
T: TryInto<i32>,
<T as TryInto<i32>>::Error: std::fmt::Debug,
{
x.try_into().expect("dimension must fit into i32")
}
#[doc(hidden)]
pub fn dims3(a_shape: impl Shape, b_shape: impl Shape, c_shape: impl Shape) -> (i32, i32, i32) {
let (m, k) = (a_shape.dim(0), a_shape.dim(1));
let (k2, n) = (b_shape.dim(0), b_shape.dim(1));
let (m2, n2) = (c_shape.dim(0), c_shape.dim(1));
assert!(m == m2, "a and c must agree in number of rows");
assert!(n == n2, "b and c must agree in number of columns");
assert!(
k == k2,
"a's number of columns must be equal to b's number of rows"
);
(into_i32(m), into_i32(n), into_i32(k))
}
#[doc(hidden)]
pub fn dims2(a_shape: impl Shape, b_shape: impl Shape) -> (i32, i32) {
let (m, k) = (a_shape.dim(0), a_shape.dim(1));
let (k2, n) = (b_shape.dim(0), b_shape.dim(1));
assert!(
k == k2,
"a's number of columns must be equal to b's number of rows"
);
(into_i32(m), into_i32(n))
}
#[doc(hidden)]
pub fn transpose_in_place<T, D0, D1, L>(c: &mut Slice<T, (D0, D1), L>)
where
T: ComplexFloat + Default,
D0: Dim,
D1: Dim,
L: Layout,
{
let (m, n) = *c.shape();
let m = m.size();
let n = n.size();
if n == m {
for i in 0..m {
for j in (i + 1)..n {
c.swap(i * n + j, j * n + i);
}
}
} else {
let mut result = tensor![[T::default(); m]; n];
for j in 0..n {
for i in 0..m {
result[j * m + i] = c[i * n + j];
}
}
for j in 0..n {
for i in 0..m {
c[j * m + i] = result[j * m + i];
}
}
}
}
#[doc(hidden)]
pub fn conjugate_in_place<T, D0, D1, L>(c: &mut Slice<T, (D0, D1), L>)
where
T: ComplexFloat + Default,
D0: Dim,
D1: Dim,
L: Layout,
{
c.iter_mut().for_each(|elem| *elem = elem.conj());
}
#[doc(hidden)]
pub fn ipiv_to_perm_mat<T: ComplexFloat, D0: Dim, D1: Dim>(
ipiv: &[i32],
m: usize,
) -> Array<T, (D0, D1)> {
let mut p = Array::from_elem(<(D0, D1) as Shape>::from_dims(&[m, m]), T::zero());
for i in 0..m {
p[[i, i]] = T::one();
}
for i in 0..ipiv.len() {
let pivot_row = (ipiv[i] - 1) as usize; if pivot_row != i {
for j in 0..m {
let temp = p[[i, j]];
p[[i, j]] = p[[pivot_row, j]];
p[[pivot_row, j]] = temp;
}
}
}
p
}
#[doc(hidden)]
pub fn to_col_major<T, D0: Dim, D1: Dim, L>(c: &Slice<T, (D0, D1), L>) -> Array<T, (D1, D0)>
where
T: ComplexFloat + Default + Clone,
L: Layout,
{
let csh = *c.shape();
let (m, n) = (csh.dim(0), csh.dim(1));
let shape = <(D1, D0) as Shape>::from_dims(&[n, m]);
let mut result = Array::<T, (D1, D0)>::zeros(shape);
for i in 0..m {
for j in 0..n {
result[[j, i]] = c[[i, j]];
}
}
result
}
pub fn trace<T, D0, D1, L>(a: &Slice<T, (D0, D1), L>) -> T
where
T: ComplexFloat + std::ops::Add<Output = T> + Copy,
D0: Dim,
D1: Dim,
L: Layout,
{
let ash = *a.shape();
let (m, n) = (ash.dim(0), ash.dim(1));
assert_eq!(m, n, "trace is only defined for square matrices");
let mut tr = T::zero();
for i in 0..n {
tr = tr + a[[i, i]];
}
tr
}
pub fn identity<T: Zero + One, D0: Dim, D1: Dim>(n: usize) -> Array<T, (D0, D1)> {
Array::<T, (D0, D1)>::from_fn(<(D0, D1) as Shape>::from_dims(&[n, n]), |i| {
if i[0] == i[1] { T::one() } else { T::zero() }
})
}
pub fn identity_k<T: Zero + One, D0: Dim, D1: Dim>(n: usize, k: isize) -> Array<T, (D0, D1)> {
Array::<T, (D0, D1)>::from_fn(<(D0, D1) as Shape>::from_dims(&[n, n]), |i| {
if (i[1] as isize - i[0] as isize) == k {
T::one()
} else {
T::zero()
}
})
}
pub fn kron<T, D0, D1, La, Lb>(
a: &Slice<T, (D0, D1), La>,
b: &Slice<T, (D0, D1), Lb>,
) -> Array<T, (D0, D1)>
where
T: ComplexFloat + std::ops::Mul<Output = T> + Copy,
D0: Dim,
D1: Dim,
La: Layout,
Lb: Layout,
{
let ash = *a.shape();
let (ma, na) = (ash.dim(0), ash.dim(1));
let bsh = *b.shape();
let (mb, nb) = (bsh.dim(0), bsh.dim(1));
let out_shape = <(D0, D1) as Shape>::from_dims(&[ma * mb, na * nb]);
Array::<T, (D0, D1)>::from_fn(out_shape, |idx| {
let i = idx[0];
let j = idx[1];
let ai = i / mb;
let bi = i % mb;
let aj = j / nb;
let bj = j % nb;
a[[ai, aj]] * b[[bi, bj]]
})
}
pub fn unravel_index<T, S: Shape, L: Layout>(x: &Slice<T, S, L>, mut flat: usize) -> Vec<usize> {
let rank = x.rank();
assert!(
flat < x.len(),
"flat index out of bounds: {} >= {}",
flat,
x.len()
);
let mut coords = vec![0usize; rank];
for i in (0..rank).rev() {
let dim = x.shape().dim(i);
coords[i] = flat % dim;
flat /= dim;
}
coords
}
pub fn diag<T: Zero + One + Clone, D: Dim>(v: &Slice<T, (D,)>) -> Array<T, (D, D)> {
let n = v.dim(0);
Array::<T, (D, D)>::from_fn(<(D, D) as Shape>::from_dims(&[n, n]), |i| {
if i[0] == i[1] {
v[i[0]].clone()
} else {
T::zero()
}
})
}